import marimo

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


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

    return mo, np, plt


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    # Extended bootstrap: average particle lifetime

    Generate $N\sim\mathrm{Poisson}(400)$ independent exponential lifetimes
    with mean $\langle t\rangle=1$. Calculate the measured mean lifetime and use
    an extended bootstrap to estimate its variance and standard deviation.
    Compare with the reference variance $1/400=0.0025$ and standard deviation $0.05$.

    For each bootstrap replica, draw independent multiplicities $w_i\sim\mathrm{Poisson}(1)$:

    $$
    \bar t=\frac1N\sum_i t_i,\qquad
    N_b=\sum_i w_i,\qquad
    \bar t_b=\frac{\sum_i w_i t_i}{N_b}.
    $$

    The variance and standard deviation of the bootstrap means estimate the
    uncertainty of the measured mean. Empty replicas are omitted.
    """)
    return


@app.cell
def _(np):
    rng = np.random.default_rng(1)
    n_observed = int(rng.poisson(400))
    if n_observed == 0:
        raise ValueError("No recorded events: the mean lifetime is undefined.")
    lifetimes = rng.exponential(scale=1, size=n_observed)
    observed_mean = float(np.mean(lifetimes))
    print(f"Recorded events: {n_observed}")
    print(f"Measured mean lifetime: {observed_mean:.6f}")
    return lifetimes, n_observed, observed_mean


@app.cell
def _(lifetimes, n_observed, np):
    _rng = np.random.default_rng(2)
    n_bootstrap = 20_000
    _means = []
    for _ in range(n_bootstrap):
        _weights = _rng.poisson(1, size=n_observed)
        _replica_count = np.sum(_weights)
        if _replica_count > 0:
            _means.append(np.sum(_weights * lifetimes) / _replica_count)
    bootstrap_means = np.array(_means)
    if len(bootstrap_means) < 2:
        raise ValueError("At least two nonempty replicas are needed.")
    bootstrap_variance = float(np.var(bootstrap_means, ddof=1))
    bootstrap_sd = float(np.sqrt(bootstrap_variance))
    print(f"Bootstrap replicas: {n_bootstrap}; empty replicas omitted: {n_bootstrap - len(bootstrap_means)}")
    return bootstrap_means, bootstrap_sd, bootstrap_variance


@app.cell
def _(bootstrap_sd, bootstrap_variance, mo):
    mo.ui.table([
        {"Method": "Extended bootstrap", "Variance of mean": bootstrap_variance,
         "Standard deviation of mean": bootstrap_sd},
        {"Method": "Exercise reference", "Variance of mean": 1 / 400,
         "Standard deviation of mean": 0.05},
    ], selection=None)
    return


@app.cell
def _(bootstrap_means, observed_mean, plt):
    _fig, _ax = plt.subplots(figsize=(9, 4), constrained_layout=True)
    _ax.hist(bootstrap_means, bins=70, density=True, alpha=0.65)
    _ax.axvline(observed_mean, color="C1", linestyle="--", label=f"Measured mean = {observed_mean:.3f}")
    _ax.axvline(1, color="k", linestyle=":", label="True mean = 1")
    _ax.set(xlabel="Bootstrap mean lifetime", ylabel="Probability density",
            title="Extended bootstrap")
    _ax.legend()
    plt.show()
    return


@app.cell
def _():
    return


@app.cell
def _():
    return


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