# Musterloesung
#
# Wir wenden den selbstgeschriebenen MH-Sampler (mh_sampler.py) auf die
# Mock-PTA-Likelihood (mock_likelihood.py) an. Als Modell benutzen wir ein
# "broken power law": ein einfaches, (noch) nicht physikalisch motiviertes
# Spektrum, das durch eine Amplitude A und eine Peak-Frequenz f_peak
# beschrieben wird:
#
#   h^2 Omega_gw(f) = A * x / (1 + x^2),   x = f / f_peak
#
# Diese Funktion hat ihr Maximum bei x = 1 (also f = f_peak) und faellt
# rechts und links davon ab -- ein "Toy model" fuer einen Peak im Spektrum.
# Logarithmiert wird daraus log10A + log10(x/(1+x^2)); wir sampeln daher
# direkt log10A und log10(f_peak) statt A und f_peak selbst.

import numpy as np
import matplotlib.pyplot as plt
from corner import corner

from mh_sampler import metropolis_hastings
from mock_likelihood import freqs, mu_log10, sigma_log10, log_likelihood, nHz


# --- Modell -------------------------------------------------------------------
def broken_powerlaw_log10(f_hz, log10A, f_pivot_hz):
    x = f_hz / f_pivot_hz
    return log10A + np.log10(x / (1.0 + x**2))


# --- Prior (gleichverteilt) ----------------------------------------------------
LOG_A_MIN, LOG_A_MAX = -13.0, -4.0
LOG_F_PEAK_MIN, LOG_F_PEAK_MAX = -10.0, -7.0   # f_peak in [0.1, 100] nHz


def log_prior(theta):
    log10A, log_f_peak = theta
    if (LOG_A_MIN <= log10A <= LOG_A_MAX) and (LOG_F_PEAK_MIN <= log_f_peak <= LOG_F_PEAK_MAX):
        return 0.0
    return -np.inf


def log_posterior(theta):
    lp = log_prior(theta)
    if not np.isfinite(lp):
        return -np.inf
    log10A, log_f_peak = theta
    model = broken_powerlaw_log10(freqs, log10A, 10**log_f_peak)
    return lp + log_likelihood(model)


# --- MCMC ------------------------------------------------------------------------
if __name__ == "__main__":
    chain, logp, acc = metropolis_hastings(
        log_posterior, x0=[-8.0, -9.0], step_sizes=[0.45, 0.22],
        n_steps=50000, seed=42)
    samples = chain[5000:]   # burn-in
    logp_samples = logp[5000:]

    print(f"Akzeptanzrate: {acc:.2f}")

    i_best = np.argmax(logp_samples)
    log10A_best, log_f_peak_best = samples[i_best]
    f_peak_best = 10**log_f_peak_best
    print(f"Bestes log10(A): {log10A_best:.3f}")
    print(f"Bestes f_peak: {f_peak_best / nHz:.2f} nHz")

    lo, med, hi = np.percentile(samples[:, 0], [16, 50, 84])
    print(f"log10(A): {med:.3f}  (+{hi-med:.3f} / -{med-lo:.3f})")
    lo, med, hi = np.percentile(samples[:, 1], [16, 50, 84])
    print(f"log10(f_peak / Hz): {med:.3f}  (+{hi-med:.3f} / -{med-lo:.3f})")

    # --- Plots -------------------------------------------------------------------
    labels = [r"$\log_{10}A$", r"$\log_{10}(f_{peak}/\mathrm{Hz})$"]
    corner(samples, labels=labels, show_titles=True, title_fmt=".3f")
    plt.savefig("fit_broken_powerlaw_2d_corner.pdf", dpi=150, bbox_inches="tight")
    plt.close()

    f_plot = np.logspace(np.log10(freqs[0]) - 1, np.log10(freqs[-1]) + 1, 500)
    model_best = broken_powerlaw_log10(f_plot, log10A_best, f_peak_best)

    plt.figure()
    plt.errorbar(freqs / nHz, mu_log10, yerr=sigma_log10, fmt="o",
                 capsize=3, label="Mock-Daten")
    plt.plot(f_plot / nHz, model_best, color="C1", label="bestes Modell")
    plt.xscale("log")
    plt.xlabel("f [nHz]")
    plt.ylabel(r"$\log_{10} h^2\Omega_{gw}$")
    plt.legend()
    plt.tight_layout()
    plt.savefig("fit_broken_powerlaw_2d_data.pdf", dpi=150, bbox_inches="tight")

    print("\nPlots gespeichert: fit_broken_powerlaw_2d_corner.pdf, "
          "fit_broken_powerlaw_2d_data.pdf")
