# Musterloesung -- "Fuer die Schnellen"
#
# Verallgemeinerung des Toy models aus fit_broken_powerlaw_2d.py:
#
#   log10(h^2 Omega_gw(f)) = A + log10(x^a / (1 + x^c)),   x = f / f_peak
#
# mit zwei zusaetzlichen freien Parametern a und c (fuer a=1, c=2 ist das
# genau das Modell aus fit_broken_powerlaw_2d.py). metropolis_hastings selbst
# bleibt dabei komplett unveraendert -- nur x0/step_sizes wachsen von 2 auf 4
# Eintraege.

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_general_log10(f_hz, A, f_peak_hz, a, c):
    x = f_hz / f_peak_hz
    return A + np.log10(x**a / (1.0 + x**c))


# --- Prior (gleichverteilt) ----------------------------------------------------
A_MIN, A_MAX = -13.0, -4.0
LOG_F_PEAK_MIN, LOG_F_PEAK_MAX = -10.0, -7.0
A_EXP_MIN, A_EXP_MAX = 0.5, 5.0
C_EXP_MIN, C_EXP_MAX = 1.0, 8.0


def log_prior(theta):
    A, log_f_peak, a, c = theta
    if not (A_MIN <= A <= A_MAX):
        return -np.inf
    if not (LOG_F_PEAK_MIN <= log_f_peak <= LOG_F_PEAK_MAX):
        return -np.inf
    if not (A_EXP_MIN <= a <= A_EXP_MAX):
        return -np.inf
    if not (C_EXP_MIN <= c <= C_EXP_MAX):
        return -np.inf
    return 0.0


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


# --- MCMC ------------------------------------------------------------------------
if __name__ == "__main__":
    chain, logp, acc = metropolis_hastings(
        log_posterior, x0=[-8.0, -9.0, 1.0, 2.0],
        step_sizes=[0.2, 0.2, 0.2, 0.35],
        n_steps=100000, seed=42)
    samples = chain[10000:]   # burn-in
    logp_samples = logp[10000:]

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

    i_best = np.argmax(logp_samples)
    A_best, log_f_peak_best, a_best, c_best = samples[i_best]
    print(f"Bestes A: {A_best:.3f}")
    print(f"Bestes f_peak: {10**log_f_peak_best / nHz:.2f} nHz")
    print(f"Bestes a: {a_best:.2f}")
    print(f"Bestes c: {c_best:.2f}")

    for name, col in zip(["log10A", "log10(f_peak/Hz)", "a", "c"], range(4)):
        lo, med, hi = np.percentile(samples[:, col], [16, 50, 84])
        print(f"{name}: {med:.3f}  (+{hi-med:.3f} / -{med-lo:.3f})")

    # --- Plots -------------------------------------------------------------------
    labels = [r"$\log_{10}A$", r"$\log_{10}(f_{peak}/\mathrm{Hz})$", r"$a$", r"$c$"]
    corner(samples, labels=labels, show_titles=True, title_fmt=".2f")
    plt.savefig("fit_broken_powerlaw_general_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_general_log10(f_plot, A_best, 10**log_f_peak_best,
                                                a_best, c_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 verallgemeinertes 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_general_data.pdf", dpi=150, bbox_inches="tight")

    print("\nPlots gespeichert: fit_broken_powerlaw_general_corner.pdf, "
          "fit_broken_powerlaw_general_data.pdf")
