# Musterloesung
#
# Allgemeiner Metropolis-Hastings (MH) Sampler.
#
# Dieses Modul stellt eine einzige Funktion `metropolis_hastings` bereit, die
# einen Random-Walk-MH-Algorithmus fuer eine beliebige (unnormierte)
# log-Posterior-Funktion in N Dimensionen durchfuehrt.
#
# Als Test wird der Sampler hier auf eine bekannte, korrelierte 2D-Gauss-
# Verteilung angewendet. So kann man pruefen, ob der eigene Sampler
# funktioniert (Mittelwert & Kovarianz der samples sollten mit den
# vorgegebenen Werten uebereinstimmen), bevor man ihn auf die (unbekannte)
# PTA-Likelihood anwendet.

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__":
    # Bekannte Zielverteilung: 2D Gauss mit Mittelwert mu und Kovarianz Sigma.
    # Da wir Mittelwert und Kovarianz hier von Hand vorgeben, koennen wir
    # spaeter die samples direkt mit der "richtigen Loesung" vergleichen.
    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)}")
    print(f"Erwarteter Mittelwert (mu): {mu}")
    print(f"Sample-Kovarianz:\n{np.cov(samples.T)}")
    print(f"Erwartete Kovarianz (Sigma):\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")
