import marimo

__generated_with = "0.24.2"
app = marimo.App()


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    # Least squares: straight-line fitting

    We fit $y=a_0+a_1x$ to 20 measurements with expectation $\mu(y)=10+10x$.
    """)
    return


@app.cell
def _():
    import marimo as mo
    import numpy as np
    import matplotlib.pyplot as plt
    from scipy.stats import chi2


    return chi2, mo, np, plt


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ## 1. Constant scatter: Gaussian, exponential and uniform

    All three errors must have zero mean and standard deviation $\sigma=4$:

    - Gaussian: $\epsilon\sim\mathcal N(0,4)$.
    - Shifted exponential: $\epsilon=E-4$, with $E\sim\mathrm{Exp}(\mathrm{scale}=4)$.
    - Uniform: $\epsilon\sim U(-4\sqrt3,4\sqrt3)$.

    The exponential must be shifted to keep $\langle y\rangle=10+10x$.


    Define $\bar x=\frac1n\sum_i x_i$, $\bar y=\frac1n\sum_i y_i$ and
    $D=\sum_i(x_i-\bar x)^2$, with $n=20$.
    Minimize the sum of squared residuals,

    $$
    S(a_0,a_1)=\sum_i(y_i-a_0-a_1x_i)^2.
    $$

    **Find the two stationary equations.**

    Differentiating with respect to the intercept and slope gives

    $$
    \frac{\partial S}{\partial a_0}=-2\sum_i(y_i-\hat a_0-\hat a_1x_i)=0,
    $$
    $$
    \frac{\partial S}{\partial a_1}=-2\sum_i x_i(y_i-\hat a_0-\hat a_1x_i)=0.
    $$

    **Express the intercept in terms of the unknown slope.**

    Expanding the first equation and dividing by $n$ gives

    $$
    \sum_i y_i-n\hat a_0-\hat a_1\sum_i x_i=0
    \quad\Longrightarrow\quad
    \hat a_0=\bar y-\hat a_1\bar x.
    $$

    **Substitute into the second equation.**

    The residual becomes

    $$
    y_i-\hat a_0-\hat a_1x_i=(y_i-\bar y)-\hat a_1(x_i-\bar x).
    $$

    Therefore,

    $$
    \sum_i x_i(y_i-\bar y)-\hat a_1\sum_i x_i(x_i-\bar x)=0,
    $$

    so

    $$
    \hat a_1=\frac{\sum_i x_i(y_i-\bar y)}{\sum_i x_i(x_i-\bar x)}.
    $$

    Using $\sum_i x_i=n\bar x$ and $\sum_i y_i=n\bar y$, the numerator is

    $$
    \sum_i x_i(y_i-\bar y)
    =\sum_i x_i y_i-n\bar x\bar y
    =\sum_i(x_i-\bar x)y_i.
    $$

    Since $\sum_i(x_i-\bar x)=0$, the denominator is

    $$
    \sum_i x_i(x_i-\bar x)
    =\sum_i(x_i-\bar x)^2
    +\bar x\underbrace{\sum_i(x_i-\bar x)}_{0}=D.
    $$

    **Calculate the slope, then the intercept.**

    $$
    \boxed{\hat a_1=\frac{\sum_i(x_i-\bar x)y_i}{D}},
    \qquad \boxed{\hat a_0=\bar y-\hat a_1\bar x}.
    $$
    """)
    return


@app.cell
def _(np):
    x = (np.arange(20) + 0.5) * 0.1
    mu = 10 + 10 * x
    truth = np.array([10., 10.])
    n_experiments = 10_000
    return mu, n_experiments, x


@app.cell
def _():
    # Approximate reference values quoted in the exercise (constant scatter σ = 4).
    reference_sigma_a0 = 1.791
    reference_sigma_a1 = 1.551
    reference_rho = -0.866
    print("Reference values from the exercise:")
    print(f"σ(a0) ≈ {reference_sigma_a0}, σ(a1) ≈ {reference_sigma_a1}, ρ ≈ {reference_rho}")
    return


@app.cell
def _(mo, mu, n_experiments, np, x):
    _rng = np.random.default_rng(1)
    constant_samples = {
        "Gaussian": mu + _rng.normal(0, 4, (n_experiments, len(x))),
        "Exponential": mu + _rng.exponential(4, (n_experiments, len(x))) - 4,
        "Uniform": mu + _rng.uniform(-4*np.sqrt(3), 4*np.sqrt(3), (n_experiments, len(x))),
    }
    constant_fits = {}
    constant_chi2 = {}
    _x_bar = np.mean(x)
    _sxx = np.sum((x - _x_bar)**2)
    for _name, _samples in constant_samples.items():
        # Slope: a1 = sum((x_i - x_bar)*(y_i)) / D,
        # where D = sum((x_i - x_bar)**2) = _sxx.
        # axis=1 sums the measurement points separately for each experiment.
        _a1 = np.sum((x - _x_bar) * _samples, axis=1) / _sxx
        # Intercept: a0 = y_bar - a1*x_bar, using the slope calculated above.
        _a0 = np.mean(_samples, axis=1) - _a1 * _x_bar
        _fits = np.column_stack([_a0, _a1])
        constant_fits[_name] = _fits
        constant_chi2[_name] = np.sum(((_samples - (_a0[:, None] + _a1[:, None] * x)) / 4)**2, axis=1)
    constant_summary = []
    for _name, _fits in constant_fits.items():
        constant_summary.append({
            "Uncertainty": _name,
            "Mean a0": float(_fits[:, 0].mean()),
            "Mean a1": float(_fits[:, 1].mean()),
            "SD a0": float(_fits[:, 0].std(ddof=1)),
            "SD a1": float(_fits[:, 1].std(ddof=1)),
            "Correlation": float(np.corrcoef(_fits.T)[0, 1]),
            "Mean chi² min": float(constant_chi2[_name].mean()),
        })
    mo.ui.table(constant_summary, selection=None)
    return constant_chi2, constant_fits, constant_samples


@app.cell
def _(constant_fits, constant_samples, mu, plt, x):
    _fig, _axes = plt.subplots(1, 3, figsize=(15, 4), constrained_layout=True)
    for _ax, (_name, _samples) in zip(_axes, constant_samples.items()):
        _ax.errorbar(x, _samples[0], yerr=4, fmt="o", label="One experiment")
        _ax.plot(x, mu, "k--", label="Expectation")
        _ax.plot(x, constant_fits[_name][0, 0] + constant_fits[_name][0, 1] * x, label="Fit")
        _ax.set(title=_name, xlabel="x", ylabel="y")
        _ax.legend()
    plt.show()
    return


@app.cell
def _(chi2, constant_chi2, constant_fits, np, plt):
    _fig, _axes = plt.subplots(1, 3, figsize=(15, 4), constrained_layout=True)
    for _name, _fits in constant_fits.items():
        for _i in range(2):
            _axes[_i].hist(_fits[:, _i], bins=70, density=True, histtype="step", label=_name)
        _axes[2].hist(constant_chi2[_name], bins=np.linspace(0, 100, 101), density=True, histtype="step", label=_name)
    for _i in range(2):
        _axes[_i].axvline(10, color="k", linestyle="--", label="True value")
        _axes[_i].set(xlabel=f"a{_i}", ylabel="Density")
    _t = np.linspace(0, 100, 500)
    _axes[2].plot(_t, chi2.pdf(_t, 18), "k--", label="χ² with 18 degrees of freedom")
    _axes[2].set(xlabel="Minimum χ² (display range 0–100)", ylabel="Density")
    for _ax in _axes:
        _ax.legend()
    plt.show()
    return


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ## 2. Estimate a common Poisson mean

    Generate repeated experiments with $y_i\sim\mathrm{Poisson}(4)$,
    using $n=10$ and $n=100$ measurements per experiment.
    Compare the distributions of these three estimates:

    **Fixed variance — arithmetic mean**

    $$
    \hat{\mu}_1=\frac{1}{n}\sum_i y_i.
    $$

    **Variance varied during the fit — quadratic mean**

    $$
    \hat{\mu}_2=\sqrt{\frac{1}{n}\sum_i y_i^2}.
    $$

    **Empirical variance — harmonic mean**

    Discard zero measurements and let $n_+$ be the number of positive counts:

    $$
    \hat{\mu}_3=\frac{n_+}{\displaystyle\sum_{i:y_i>0}\frac{1}{y_i}}.
    $$

    All three methods use the same generated experiments.
    """)
    return


@app.cell
def _(mo, n_experiments, np):
    _rng = np.random.default_rng(3)
    poisson_mean_estimates = {}
    _rows = []
    for _n in (10, 100):
        _counts = _rng.poisson(4, size=(n_experiments, _n))
        _positive = _counts > 0
        _n_positive = np.sum(_positive, axis=1)
        _inverse = np.divide(
            1.0, _counts, out=np.zeros_like(_counts, dtype=float), where=_positive
        )
        _valid = _n_positive > 0
        _estimates = {
            "Arithmetic mean": np.mean(_counts, axis=1),
            "Quadratic mean": np.sqrt(np.mean(_counts**2, axis=1)),
            "Harmonic mean": _n_positive[_valid] / np.sum(_inverse[_valid], axis=1),
        }
        poisson_mean_estimates[_n] = _estimates
        for _method, _values in _estimates.items():
            _rows.append({
                "n": _n, "Estimator": _method,
                "Mean estimate (true value: 4)": float(np.mean(_values)),
                "Standard deviation": float(np.std(_values, ddof=1)),
            })
        if np.any(~_valid):
            print(f"n={_n}: {np.count_nonzero(~_valid)} all-zero experiments omitted from the harmonic mean.")
    mo.ui.table(_rows, selection=None)
    return (poisson_mean_estimates,)


@app.cell
def _(np, plt, poisson_mean_estimates):
    _fig, _axes = plt.subplots(2, 3, figsize=(15, 8), constrained_layout=True)
    _bins = np.linspace(0, 8, 81)
    for _row, (_n, _estimates) in enumerate(poisson_mean_estimates.items()):
        for _col, (_method, _values) in enumerate(_estimates.items()):
            _ax = _axes[_row, _col]
            _ax.hist(_values, bins=_bins, density=True, alpha=0.65, color=f"C{_col}")
            _ax.axvline(4, color="k", linestyle="--", label="True mean = 4")
            _ax.set(title=f"{_method}, n = {_n}",
                    xlabel="Estimated Poisson mean", ylabel="Density", xlim=(0, 8))
            _ax.legend()
    plt.show()
    return


@app.cell(hide_code=True)
def _(mo):
    mo.md("""
    The arithmetic mean is centered near 4. The quadratic mean overestimates it,
    while the harmonic mean underestimates it. Increasing the sample size narrows
    the distributions but does not remove the biases of the latter two methods.
    """)
    return


@app.cell
def _():
    return


if __name__ == "__main__":
    app.run()
