# Musterloesung
#
# Allgemeiner Metropolis-Hastings (MH) Sampler -- identisch zu dem, den wir
# schon im Inselpendel-Arbeitsblatt gebaut und getestet haben (siehe
# Aufgaben/DSA_2026/Inselpendel/mh_sampler.py). Das ist die ganze Pointe eines
# allgemeinen MCMC-Samplers: er ist "dimensionsblind" und weiss nichts von
# Pendeln oder Gravitationswellen -- nur `log_post` aendert sich.
#
# Als Selbsttest wird der Sampler hier wieder auf eine bekannte, korrelierte
# 2D-Gauss-Verteilung angewendet (Mittelwert & Kovarianz der samples sollten
# mit den vorgegebenen Werten uebereinstimmen), bevor wir ihn auf die
# (unbekannte) PTA-Likelihood loslassen.

import numpy as np
import matplotlib.pyplot as plt


def metropolis_hastings(log_post, x0, step_sizes, n_steps, seed=0):
    """
    Random-Walk Metropolis-Hastings Sampler.

    Parameters
    ----------
    log_post : callable
        Funktion theta -> log(posterior(theta)) (bis auf eine additive
        Konstante). Soll -np.inf zurueckgeben, wenn theta ausserhalb des
        erlaubten Bereichs (Prior) liegt.
    x0 : array_like, shape (ndim,)
        Startpunkt im Parameterraum.
    step_sizes : array_like, shape (ndim,)
        Standardabweichungen der (Gauss'schen) Vorschlagsverteilung, eine
        pro Parameter.
    n_steps : int
        Anzahl der MCMC-Schritte.
    seed : int
        Seed des Zufallszahlengenerators (fuer Reproduzierbarkeit).

    Returns
    -------
    chain : ndarray, shape (n_steps, ndim)
        Die gesampelten Parameterwerte.
    log_post_chain : ndarray, shape (n_steps,)
        log-Posterior-Werte entlang der chain.
    acceptance_rate : float
        Anteil der akzeptierten Vorschlaege (Richtwert: ca. 0.2 - 0.5).
    """
    rng = np.random.default_rng(seed)

    ndim = len(x0)
    x_curr = np.array(x0, dtype=float)
    logp_curr = log_post(x_curr)

    chain = np.zeros((n_steps, ndim))
    log_post_chain = np.zeros(n_steps)
    n_accept = 0

    for i in range(n_steps):
        # 1) Vorschlag: kleiner zufaelliger Schritt von der aktuellen Position
        x_prop = x_curr + rng.normal(0.0, step_sizes, size=ndim)

        # 2) "Guete" des Vorschlags bewerten
        logp_prop = log_post(x_prop)

        # 3) Metropolis-Kriterium: Annahme-Wahrscheinlichkeit
        #    min(1, exp(logp_prop - logp_curr))
        log_alpha = logp_prop - logp_curr

        # 4) Vorschlag annehmen oder verwerfen
        if np.isfinite(logp_prop) and np.log(rng.random()) < log_alpha:
            x_curr, logp_curr = x_prop, logp_prop
            n_accept += 1

        chain[i] = x_curr
        log_post_chain[i] = logp_curr

    return chain, log_post_chain, n_accept / n_steps


# -----------------------------------------------------------------------------
# Beispiel / Selbsttest: Sampling einer bekannten, korrelierten 2D-Gaussverteilung
# -----------------------------------------------------------------------------
if __name__ == "__main__":
    mu = np.array([1.0, -0.5])
    Sigma = np.array([[1.0, 0.6],
                       [0.6, 0.5]])
    Sigma_inv = np.linalg.inv(Sigma)

    def log_post_gauss(theta):
        d = theta - mu
        return -0.5 * d @ Sigma_inv @ d

    n_steps = 20000
    burn_in = 2000

    chain, logp, acc = metropolis_hastings(
        log_post_gauss, x0=[0.0, 0.0], step_sizes=[0.7, 0.7],
        n_steps=n_steps, seed=1)

    samples = chain[burn_in:]

    print(f"Akzeptanzrate: {acc:.2f}  (sollte ungefaehr zwischen 0.2 und 0.5 liegen)")
    print(f"Sample-Mittelwert: {samples.mean(axis=0)}  (Soll: {mu})")
    print(f"Sample-Kovarianz:\n{np.cov(samples.T)}\n(Soll:\n{Sigma})")

    fig, axes = plt.subplots(1, 2, figsize=(10, 4))

    axes[0].plot(chain[:, 0], lw=0.5)
    axes[0].axvline(burn_in, color="k", ls="--", label="Ende burn-in")
    axes[0].set_xlabel("Schritt")
    axes[0].set_ylabel(r"$\theta_1$")
    axes[0].set_title("Trace-plot")
    axes[0].legend()

    axes[1].plot(samples[:, 0], samples[:, 1], '.', ms=1, alpha=0.3)
    axes[1].set_xlabel(r"$\theta_1$")
    axes[1].set_ylabel(r"$\theta_2$")
    axes[1].set_title("Samples nach dem burn-in")

    fig.tight_layout()
    fig.savefig("mh_sampler_test.pdf", dpi=150, bbox_inches="tight")
    print("\nPlot gespeichert als mh_sampler_test.pdf")
