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
    from scipy.stats import norm, truncnorm

    return mo, norm, np, plt, truncnorm


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    # Exercise 4: signal and background yields with COWs

    Generate a Gaussian peak centered at zero with standard deviation $\sigma=0.05$
    on a uniform background in $x\in[-1,1]$. Choose a signal fraction $s$ and use
    the COWs for $I(x)=1$ to estimate the two component yields and their covariance:

    $$
    w_0(x)=\frac{\sqrt{2}\,e^{-x^2/(2\sigma^2)}-\sqrt{\pi}\sigma}
    {1-\sqrt{\pi}\sigma},
    $$
    $$
    w_1(x)=\frac{1-\sqrt{2}\,e^{-x^2/(2\sigma^2)}}
    {1-\sqrt{\pi}\sigma}.
    $$

    Here 0 denotes signal and 1 background. Sum the weights to estimate the yields:

    $$
    \hat N_0=\sum_i w_0(x_i),\qquad \hat N_1=\sum_i w_1(x_i).
    $$

    The covariance estimate for a Poisson event count is

    $$
    \hat C_{00}=\sum_i w_0(x_i)^2,\qquad
    \hat C_{11}=\sum_i w_1(x_i)^2,\qquad
    \hat C_{01}=\hat C_{10}=\sum_i w_0(x_i)w_1(x_i).
    $$

    Choose $s$ with the slider (default 0.5) and draw the event count with expectation $\langle N\rangle=10000$
    to match this Poisson-count covariance convention.
    """)
    return


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ## Analytical construction

    The normalized component densities are, to negligible Gaussian-tail corrections,

    $$
    g_0(x)=\frac{e^{-x^2/(2\sigma^2)}}{\sqrt{2\pi}\sigma},\qquad g_1(x)=\frac12.
    $$

    Write $w_0(x)=a g_0(x)+b g_1(x)$. To count signal and cancel background, require

    $$
    \int_{-1}^1w_0(x)g_0(x)\,dx=1,\qquad
    \int_{-1}^1w_0(x)g_1(x)\,dx=0.
    $$

    With $\sigma=0.05$, the boundaries are 20 standard deviations from the peak.
    Extending the Gaussian integrals to the whole real line gives

    $$
    \int_{-1}^1g_0^2\,dx\simeq\frac{1}{2\sqrt\pi\sigma},\qquad
    \int_{-1}^1g_0g_1\,dx\simeq\frac12,\qquad
    \int_{-1}^1g_1^2\,dx=\frac12.
    $$

    The two conditions become

    $$
    \frac{a}{2\sqrt\pi\sigma}+\frac b2=1,\qquad
    \frac a2+\frac b2=0.
    $$

    Thus $b=-a$ and

    $$
    a=\frac{2\sqrt\pi\sigma}{1-\sqrt\pi\sigma},\qquad
    b=-\frac{2\sqrt\pi\sigma}{1-\sqrt\pi\sigma}.
    $$

    Substituting the densities gives

    $$
    w_0(x)=\frac{\sqrt2 e^{-x^2/(2\sigma^2)}-\sqrt\pi\sigma}{1-\sqrt\pi\sigma}.
    $$

    For the background, write $w_1(x)=c g_0(x)+d g_1(x)$ and reverse the conditions:

    $$
    \frac{c}{2\sqrt\pi\sigma}+\frac d2=0,\qquad
    \frac c2+\frac d2=1.
    $$

    This gives

    $$
    c=-\frac{2\sqrt\pi\sigma}{1-\sqrt\pi\sigma},\qquad
    d=\frac{2}{1-\sqrt\pi\sigma},
    $$
    $$
    w_1(x)=\frac{1-\sqrt2e^{-x^2/(2\sigma^2)}}{1-\sqrt\pi\sigma}.
    $$
    """)
    return


@app.cell(hide_code=True)
def _(mo):
    signal_fraction_slider = mo.ui.slider(
        0, 1, step=0.01, value=0.5,
        label="Signal fraction s", show_value=True, debounce=True,
    )
    signal_fraction_slider
    return (signal_fraction_slider,)


@app.cell
def _(np, signal_fraction_slider, truncnorm):
    sigma = 0.05
    signal_fraction = signal_fraction_slider.value
    rng = np.random.default_rng(1)
    n_events = int(rng.poisson(10_000))
    _is_signal = rng.random(n_events) < signal_fraction
    actual_signal = int(np.sum(_is_signal))
    actual_background = n_events - actual_signal
    data = np.empty(n_events)
    data[_is_signal] = truncnorm.rvs(
        -1 / sigma, 1 / sigma, loc=0, scale=sigma,
        size=actual_signal, random_state=rng,
    )
    data[~_is_signal] = rng.uniform(-1, 1, size=actual_background)
    print(n_events)
    return actual_background, actual_signal, data, n_events, sigma


@app.cell
def _(actual_background, actual_signal, data, mo, np, sigma):
    # Apply the supplied COWs to every event, without using its simulated label.
    _peak = np.sqrt(2) * np.exp(-data**2 / (2 * sigma**2))
    _denominator = 1 - np.sqrt(np.pi) * sigma
    w0 = (_peak - np.sqrt(np.pi) * sigma) / _denominator
    w1 = (1 - _peak) / _denominator
    signal_yield = float(np.sum(w0))
    background_yield = float(np.sum(w1))

    # Yield variances and covariance: sums of squared weights and cross-products.
    C00 = float(np.sum(w0**2))
    C11 = float(np.sum(w1**2))
    C01 = float(np.sum(w0 * w1))
    covariance = np.array([[C00, C01], [C01, C11]])
    correlation = C01 / np.sqrt(C00 * C11)
    print("Covariance matrix (signal, background):")
    print(covariance)
    print(f"Correlation: {correlation:.4f}")
    mo.ui.table([
        {"Component": "Signal", "Generated count": actual_signal,
         "Estimated yield": signal_yield, "Standard deviation": np.sqrt(C00)},
        {"Component": "Background", "Generated count": actual_background,
         "Estimated yield": background_yield, "Standard deviation": np.sqrt(C11)},
    ], selection=None)
    return background_yield, signal_yield


@app.cell
def _(background_yield, data, n_events, norm, np, plt, sigma, signal_yield):
    _edges = np.linspace(-1, 1, 101)
    _counts, _edges = np.histogram(data, bins=_edges)
    _centers = (_edges[:-1] + _edges[1:]) / 2
    # Integrate each component over the bins to compare expected counts with data.
    _signal_bins = signal_yield * np.diff(norm.cdf(_edges, scale=sigma))
    _background_bins = background_yield * np.diff(_edges) / 2
    _fig, _ax = plt.subplots(figsize=(10, 5), constrained_layout=True)
    _ax.errorbar(_centers, _counts, yerr=np.sqrt(_counts), fmt=".", color="black", label="Data")
    _ax.stairs(_signal_bins + _background_bins, _edges, label="COW yield estimate: total")
    _ax.stairs(_signal_bins, _edges, linestyle="--", label="Signal")
    _ax.stairs(_background_bins, _edges, linestyle=":", label="Background")
    _ax.set(xlabel="x", ylabel="Events per bin", title=f"Gaussian signal and uniform background: N = {n_events}")
    _ax.legend()
    plt.show()
    return


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    Since $w_0(x)+w_1(x)=1$, the estimated yields sum to the observed event count.
    """)
    return


@app.cell
def _():
    return


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