import marimo

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


@app.cell
def _():
    import marimo as mo

    return (mo,)


@app.cell
def _():
    import matplotlib.pyplot as plt
    import numpy as np
    #from resample.bootstrap import variance

    # create instance of pseudo-random number generator with fixed seed = 1
    rng = np.random.default_rng(seed=1)
    return np, plt, rng


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

    ## Task 1: Expectation values of transformed variables

    Assume measurements of  $\hat{x}$ ~ $\mathcal{G}(1,\sigma)$.

    Now conside $y = e^x$. Determine, with $10^6$ value of $\hat{x}$ per configuration, the bias in the following cases:

    - ${ \langle \hat{y} \rangle} - e$ as function of $ \sigma \in [0,1]$ (no correction)
    - with leading order correction with derivatives taken at true value $x=1$ ($B = e\sigma^2/2$)
    - with leading order correction with derivatives taken at measured values $\hat{x}$ ($B = e^{\hat{x}}\sigma^2/2$)

    ---
    """)
    return


@app.cell
def _(np, plt, rng):
    # Number of measurements for each sigma
    _n_samples = 1000000
    sigmas = np.linspace(0, 1, 51)
    # True values
    _x_true = 1.0
    _y_true = np.exp(_x_true)

    bias_no_correction = []
    bias_true_correction = []
    bias_measured_correction = []
    # Store biases
    for _sigma in sigmas:
        x_hat = rng.normal(loc=_x_true, scale=_sigma, size=_n_samples)
        y_hat = np.exp(x_hat)
        bias_no_correction.append(np.mean(y_hat) - _y_true)
        B_true = np.exp(_x_true) * _sigma ** 2 / 2
        y_hat_corr_true = y_hat - B_true
        bias_true_correction.append(np.mean(y_hat_corr_true) - _y_true)
        B_measured = np.exp(x_hat) * _sigma ** 2 / 2  # Generate measurements of x
        y_hat_corr_measured = y_hat - B_measured
        bias_measured_correction.append(np.mean(y_hat_corr_measured) - _y_true)
    bias_no_correction = np.array(bias_no_correction)  # Transform to y
    bias_true_correction = np.array(bias_true_correction)
    bias_measured_correction = np.array(bias_measured_correction)
    plt.figure(figsize=(8, 6))  

    plt.plot(sigmas, bias_no_correction, marker='o', markersize=3, label='No correction') 
    plt.plot(sigmas, bias_true_correction, marker='o', markersize=3, label='Correction at $x=1$') 
    plt.plot(sigmas, bias_measured_correction, marker='o', markersize=3, label='Correction at $\\hat{x}$')
    plt.axhline(0, linestyle='--')
    plt.xlabel('$\\sigma$')
    plt.ylabel('$\\langle \\hat{y} \\rangle - e$')  
    plt.legend()  
    plt.grid(alpha=0.3)  # Convert to numpy arrays


    plt.show()  
    return


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

    ### Task 2: Covariance matrix transformation

    Consider $x$ and $y$, two uncorrelated random variables normally distributed with $\langle x \rangle = 2$ and $\langle y \rangle = 1$ and $\sigma=0.1$.

    $p=x \cdot y$ and $q=x/y$ have true covariances $C_{pp}=0.0501$, $C_{pq}=-0.0313$ and $C_{qq}=0.0538$.

    Generate 1.000.000 pairs of measurements and use error propagation to determine for each pair the estimated covariance matrix elements. Compare the distributions of those estimates to the respective true values.

    ---
    """)
    return


@app.cell
def _(np, plt):
    _n_samples = 1000000
    _x_true = 2.0
    _y_true = 1.0
    _sigma_x = 0.1
    _sigma_y = 0.1
    _rng = np.random.default_rng(seed=1)
    Cpp_true = 0.0501
    Cpq_true = -0.0313
    Cqq_true = 0.0538
    x_sample = _rng.normal(loc=_x_true, scale=_sigma_x, size=_n_samples)
    y_sample = _rng.normal(loc=_y_true, scale=_sigma_y, size=_n_samples)
    p_sample = x_sample * y_sample
    q_sample = x_sample / y_sample

    # ------------------------------------------------------------
    # Generate measurements
    cov_sample = np.cov(p_sample, q_sample)
    print('Covariance matrix from generated p and q:')
    print(cov_sample)
    print('\nComparison with reference values:')
    print(f'Cpp: sample = {cov_sample[0, 0]:.4f}, reference = {Cpp_true:.4f}')
    print(f'Cpq: sample = {cov_sample[0, 1]:.4f}, reference = {Cpq_true:.4f}')
    print(f'Cqq: sample = {cov_sample[1, 1]:.4f}, reference = {Cqq_true:.4f}')
    dp_dx = y_sample
    # Calculate p and q and their empirical covariance matrix
    dp_dy = x_sample
    dq_dx = 1 / y_sample
    dq_dy = -x_sample / y_sample ** 2
    Cpp_est = _sigma_x ** 2 * dp_dx ** 2 + _sigma_y ** 2 * dp_dy ** 2
    Cpq_est = _sigma_x ** 2 * dp_dx * dq_dx + _sigma_y ** 2 * dp_dy * dq_dy
    Cqq_est = _sigma_x ** 2 * dq_dx ** 2 + _sigma_y ** 2 * dq_dy ** 2
    fig, axes = plt.subplots(3, 1, figsize=(8, 12), constrained_layout=True)
    estimates = [Cpp_est, Cpq_est, Cqq_est]
    true_values = [Cpp_true, Cpq_true, Cqq_true]
    labels = ['$C_{pp}$', '$C_{pq}$', '$C_{qq}$']
    for ax, estimate, true_value, label in zip(axes, estimates, true_values, labels):
        ax.hist(estimate, bins=100, density=True, alpha=0.5, label='Estimated')
        ax.axvline(true_value, linestyle='--', linewidth=2, label=f'True = {true_value:.4f}')
        ax.axvline(np.mean(estimate), linestyle=':', linewidth=2, label=f'Mean estimate = {np.mean(estimate):.4f}')
        ax.set_xlabel(label)
        ax.set_ylabel('Density')
        ax.set_title(f'{label}: $\\langle \\hat{{C}} \\rangle={np.mean(estimate):.4f}$, $\\sigma={np.std(estimate):.4f}$')

        ax.legend()

        ax.grid(alpha=0.3)

    # Plot distributions
    plt.show()
    return


@app.cell(hide_code=True)
def _(mo):
    covariance_mu_x = mo.ui.slider(-5, 5, step=0.1, value=2.0, label="mu_x", show_value=True, debounce=True)
    covariance_mu_y = mo.ui.slider(-5, 5, step=0.1, value=1.0, label="mu_y", show_value=True, debounce=True)
    covariance_sigma_x = mo.ui.slider(0.01, 2, step=0.01, value=0.1, label="sigma_x", show_value=True, debounce=True)
    covariance_sigma_y = mo.ui.slider(0.01, 2, step=0.01, value=0.1, label="sigma_y", show_value=True, debounce=True)
    mo.vstack([
        mo.md("Explore x, y, p = xy, q = x/y and 1/y — plots update when you release a slider."),
        mo.hstack([covariance_mu_x, covariance_sigma_x]),
        mo.hstack([covariance_mu_y, covariance_sigma_y]),
        mo.md("Ratios can have extreme values when y is close to zero."),
    ])
    return (
        covariance_mu_x,
        covariance_mu_y,
        covariance_sigma_x,
        covariance_sigma_y,
    )


@app.cell(hide_code=True)
def _(mo):
    distribution_log_x = mo.ui.checkbox(label="Log horizontal axis (symmetric log)")
    distribution_log_y = mo.ui.checkbox(label="Log density axis")
    distribution_trim_tails = mo.ui.checkbox(value=False, label="Trim 1/y tails (same mask for all distributions)")
    distribution_trim_percent = mo.ui.slider(0, 1, step=0.01, value=0.1, label="Trim percentage per tail (%)", show_value=True, debounce=True)
    mo.vstack([
        mo.hstack([distribution_log_x, distribution_log_y]),
        mo.hstack([distribution_trim_tails, distribution_trim_percent]),
        mo.md("The horizontal log scale supports negative values and is linear near zero (−0.01 to 0.01)."),
    ])
    return (
        distribution_log_x,
        distribution_log_y,
        distribution_trim_percent,
        distribution_trim_tails,
    )


@app.cell
def _(
    covariance_mu_x,
    covariance_mu_y,
    covariance_sigma_x,
    covariance_sigma_y,
    distribution_log_x,
    distribution_log_y,
    distribution_trim_percent,
    distribution_trim_tails,
    np,
    plt,
):
    _n_samples = 1000000
    _x_true = covariance_mu_x.value
    _y_true = covariance_mu_y.value
    _sigma_x = covariance_sigma_x.value
    _sigma_y = covariance_sigma_y.value
    _rng = np.random.default_rng(seed=1)
    _x_sample = _rng.normal(loc=_x_true, scale=_sigma_x, size=_n_samples)
    _y_sample = _rng.normal(loc=_y_true, scale=_sigma_y, size=_n_samples)
    _p_sample = _x_sample * _y_sample
    _q_sample = _x_sample / _y_sample
    _inverse_y_sample = 1 / _y_sample
    _trim_percent = distribution_trim_percent.value
    _trim_tails = distribution_trim_tails.value and _trim_percent > 0
    if _trim_tails:
        _lower, _upper = np.percentile(
            _inverse_y_sample, [_trim_percent, 100 - _trim_percent]
        )
        _mask = (_inverse_y_sample >= _lower) & (_inverse_y_sample <= _upper)
        _x_sample = _x_sample[_mask]
        _y_sample = _y_sample[_mask]
        _p_sample = _p_sample[_mask]
        _q_sample = _q_sample[_mask]
        _inverse_y_sample = _inverse_y_sample[_mask]

    # Distributions of the inputs and transformed variables, using the same draws.
    _distributions = [
        (r"$x$", _x_sample),
        (r"$y$", _y_sample),
        (r"$p = xy$", _p_sample),
        (r"$q = x/y$", _q_sample),
        (r"$1/y$", _inverse_y_sample),
    ]
    _fig, _axes = plt.subplots(3, 2, figsize=(12, 12), constrained_layout=True)
    if _trim_tails:
        _fig.suptitle(f"Lowest and highest {_trim_percent:g}% of 1/y removed; same sample mask applied to all")
    for _ax, (_label, _values) in zip(_axes.flat, _distributions):
        _mean = np.mean(_values)
        _std = np.std(_values)
        _ax.hist(_values, bins=100, density=True, alpha=0.65)
        _ax.axvline(_mean, color="C1", linestyle="--", label=f"Mean = {_mean:.4f}")
        _ax.set_title(f"{_label}: Std = {_std:.4f}")
        _ax.set_xlabel(_label)
        _ax.set_ylabel("Density")
        if distribution_log_x.value:
            _ax.set_xscale("symlog", linthresh=0.01)
        if distribution_log_y.value:
            _ax.set_yscale("log")
        _ax.legend()
        _ax.grid(alpha=0.3)
    _axes.flat[-1].set_visible(False)
    plt.show()
    return


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


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ---
    ### Task 3
    Generate $10^5$ n-tuples of uniformly distributed random numbers $x_i$ determine numerically, as a function of the size N of the n-tuples, the standard deviation for two estimators a and b of the central point:
    - $a = \frac{1}{N} \sum_{i=1}^{N}x_i$
    - $b = \frac{1}{2} \left ( \min_{i=1}^{N} x_i + \max_{i=1}^{N} x_i \right )$

    Hint: both are unbiased so you can compute RMS difference between estimator and $1/2$.

    ---
    """)
    return


@app.cell
def _(np, plt, rng):
    _n_samples = 100000
    N_values = np.arange(2, 101)
    sigma_a = []
    sigma_b = []
    for N in N_values:
        _x = rng.uniform(low=0, high=1, size=(_n_samples, N))
        a = np.mean(_x, axis=1)
        b = (np.min(_x, axis=1) + np.max(_x, axis=1)) / 2

    # Generate n-tuples and calculate estimators
        sigma_a.append(np.sqrt(np.mean((a - 0.5) ** 2)))
        sigma_b.append(np.sqrt(np.mean((b - 0.5) ** 2)))
        print(f'N = {N:3d} | sigma_a = {sigma_a[-1]:.6f} | sigma_b = {sigma_b[-1]:.6f}')
    sigma_a = np.array(sigma_a)
    sigma_b = np.array(sigma_b)
    sigma_a_theory = 1 / np.sqrt(12 * N_values)
    sigma_b_theory = 1 / np.sqrt(2 * (N_values + 1) * (N_values + 2))  
    plt.figure(figsize=(10, 6))
    edges = np.arange(1.5, 101.5)
    hist_a = plt.stairs(sigma_a, edges, label='$\\sigma_a$ simulation')  
    hist_b = plt.stairs(sigma_b, edges, label='$\\sigma_b$ simulation')
    plt.loglog(N_values, sigma_a_theory, '--', color=hist_a.get_edgecolor(), linewidth=2, label='$1/\\sqrt{12N}$')
    plt.loglog(N_values, sigma_b_theory, '--', color=hist_b.get_edgecolor(), linewidth=2, label='$1/\\sqrt{2(N+1)(N+2)}$')
    plt.xlabel('$N$')
    plt.ylabel('Standard deviation')
    plt.legend()
    plt.grid(alpha=0.3)
    # Plot
    plt.show()
    return


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ---
    ### Task 4
    Determine numerically the coverage probability of the interval $[n-\sqrt{n},n+\sqrt{n}]$, where $n$ is a Poisson distributed random variable, for Poisson averages $\mu=10$ and $\mu=100$.

    Determine numerically the coverage probability of the interval $[x-1,x+1]$, where $x$ is a Gaussian distributed random variable with unit variance, for the mean values $\mu=0$, $\mu=1$ and $\mu=2$.

    ---
    """)
    return


@app.cell
def _(np, rng):
    # ------------------------------------------------------------
    # Poisson coverage
    _n_samples = 100000
    _mu_values = [10, 100, 100]
    print('Poisson coverage')
    for _mu in _mu_values:
        n = rng.poisson(lam=_mu, size=_n_samples)
        lower = n - np.sqrt(n)
        upper = n + np.sqrt(n)
        _covered = (lower <= _mu) & (_mu <= upper)
        _coverage = np.mean(_covered)
        print(f'mu = {_mu:3d} | coverage = {_coverage:.4f}')
    print('')
    _n_samples = 100000
    _mu_values = [0, 1, 2]
    print('Gaussian coverage')
    for _mu in _mu_values:
        _x = rng.normal(loc=_mu, scale=1, size=_n_samples)
        lower = _x - 1
        upper = _x + 1
        _covered = (lower <= _mu) & (_mu <= upper)
        _coverage = np.mean(_covered)
        print(f'mu = {_mu:3d} | coverage = {_coverage:.4f}')
    # Gaussian coverage
    print('')
    return


@app.cell(hide_code=True)
def _(mo):
    mo.md(r"""
    ---
    ### Task 5
    In the determination of a parameter $\mu$, individual measurements are Gaussian distributed around $\mu$ with unit variance. Determine numerically, for some values in the range $\mu \in [0,10]$, the coverage probabilities of the flip-flopping 90% confidence level regions.

    ---
    """)
    return


@app.cell
def _(np, plt, rng):
    _n_samples = 100000
    _mu_values = np.linspace(0, 10, 201)
    z_one_sided = 1.28155
    z_two_sided = 1.64485
    _coverage = []
    for _mu in _mu_values:
        _x = rng.normal(loc=_mu, scale=1, size=_n_samples)
        _covered = np.zeros(_n_samples, dtype=bool)
        mask_negative = _x < 0
        mask_upper = (_x >= 0) & (_x < 3)
    # ------------------------------------------------------------
    # Calculate coverage
        mask_two_sided = _x >= 3
        _covered[mask_negative] = _mu <= z_one_sided
        _covered[mask_upper] = _mu <= _x[mask_upper] + z_one_sided
        _covered[mask_two_sided] = (_x[mask_two_sided] - z_two_sided <= _mu) & (_mu <= _x[mask_two_sided] + z_two_sided)
        _coverage.append(np.mean(_covered))
    _coverage = np.array(_coverage)
    plt.figure(figsize=(10, 6))
    plt.plot(_mu_values, _coverage, linewidth=2, label='Flip-flopping coverage')  # --------------------------------------------------------
    plt.axhline(0.9, linestyle='--', label='Nominal 90% CL')  # Flip-flopping confidence region
    plt.xlabel('$\\mu$')  #
    plt.ylabel('Coverage probability')  # x < 0:
    plt.xlim(0, 10)  #     mu < 1.28155
    plt.ylim(0.7, 1.02)  #
    plt.legend()  # 0 <= x < 3:
    plt.grid(alpha=0.3)  #     mu < x + 1.28155
    # Plot
    plt.show()  #  # x >= 3:  #     x - 1.64485 <= mu <= x + 1.64485  # --------------------------------------------------------
    return


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