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"""
    # MCMC: distribution of z

    Determine the distribution of $z=\sqrt{x^2+y^2}$ for independent positive
    variables with densities proportional to

    $$
    f(x)=\cos^2(x)e^{-x},\qquad
    g(y)=\frac{1}{1+(y-1)^2},\qquad x,y>0.
    $$

    Use a Metropolis random walk in $u=\log x$ and $v=\log y$.
    This keeps both variables positive. The target density in these coordinates
    includes the Jacobian $xy$:

    $$
    \log\rho(u,v)=2\log|\cos(e^u)|-e^u
    -\log\left[1+(e^v-1)^2\right]+u+v.
    $$

    Accept a proposal with probability $\min(1,\rho_{\mathrm{new}}/\rho_{\mathrm{old}})$.
    If rejected, keep the current state. Discard the initial burn-in, then calculate
    $z$ for each retained state and plot its histogram.
    """)
    return


@app.cell
def _(np):
    def log_density(u, v):
        x, y = np.exp(u), np.exp(v)
        return 2*np.log(abs(np.cos(x))) - x - np.log1p((y - 1)**2) + u + v


    return (log_density,)


@app.cell
def _(log_density, np):
    rng = np.random.default_rng(142)
    n_steps = 100_000
    burn_in = 1_000
    step_size = np.array([1.0, 1.3])
    current = np.log([1.0, 1.0])
    current_log_density = log_density(*current)
    z_samples = np.empty(n_steps - burn_in)

    for _step in range(n_steps):
        _proposal = current + rng.normal(size=2) * step_size
        _proposal_log_density = log_density(*_proposal)
        # Symmetric proposals in log coordinates: compare the target densities.
        if np.log(rng.random()) < min(0.0, _proposal_log_density - current_log_density):
            current = _proposal
            current_log_density = _proposal_log_density
        # Record every state, including repeated states after rejection.
        if _step >= burn_in:
            _x, _y = np.exp(current)
            z_samples[_step - burn_in] = np.hypot(_x, _y)
    return (z_samples,)


@app.cell
def _(np, plt, z_samples):
    _fig, _ax = plt.subplots(figsize=(9, 5), constrained_layout=True)
    # Show the main part of the distribution; retain all samples in normalization.
    _counts, _edges = np.histogram(z_samples, bins=np.linspace(0, 12, 61))
    _density = _counts / (len(z_samples) * np.diff(_edges))
    _ax.stairs(_density, _edges, fill=True, alpha=0.65)
    _ax.set(xlabel=r"$z=\sqrt{x^2+y^2}$", ylabel="Probability density",
            title="Distribution of z from MCMC")
    _ax.grid(alpha=0.2)
    plt.show()
    print(f"Shown range: 0–12; {100*np.mean(z_samples > 12):.2f}% of samples lie above 12.")
    return


@app.cell
def _():
    return


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