"""
run_advanced_experiments.py
===========================
Advanced characterisation experiments from the paper.

Experiment 1: Hybrid Signal Crossover
    Signal = (1-gamma)*periodic + gamma*NARMA-50
    gamma sweeps 0.0 -> 1.0 in steps of 0.15
    Shows K-R improvement grows continuously with aperiodic fraction.

Experiment 2: Noise Robustness Stress Test
    Add AWGN to NARMA-50 inputs at SNR = 30, 20, 15, 10, 5, 0, -5 dB
    Compare ESN vs K-R degradation.

Run:
    python run_advanced_experiments.py

Saves: results_advanced.pkl, figH_crossover.png, figI_noise.png
"""

import numpy as np
import itertools
import pickle
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
from scipy import stats

from benchmarks import narma
from kr_reservoir import KRReservoir
from esn_baseline import StandardESN, grid_search

# ── Config ────────────────────────────────────────────────────
N        = 200
N_SEEDS  = 10
WASHOUT  = 200
TEST_SZ  = 500
VAL_SEED = 99
M_HYB    = 50
AS_HYB   = 3.0 / M_HYB

RHO_GRID = [0.5, 0.7, 0.9, 0.95, 0.99]
LAM_GRID = [1e-6, 1e-4, 1e-2, 1.0, 10.0]

KF = 1.5 * (200 / N) ** 0.40
KS = 1.0 * (200 / N) ** 0.40


# ── Data generators ───────────────────────────────────────────
def make_periodic(T: int, seed: int = 0):
    """Quasi-periodic signal: sum of sine waves + small noise."""
    np.random.seed(seed)
    t = np.linspace(0, 20 * np.pi, T + 1)
    s = (np.sin(t)
         + 0.5 * np.sin(2.3 * t)
         + 0.3 * np.sin(0.7 * t)
         + 0.1 * np.random.randn(T + 1))
    s = (s - s.mean()) / (s.std() + 1e-12)
    return s[:-1], s[1:]   # (input, target)


def make_hybrid(gamma: float, T: int, seed: int = 0):
    """Mix periodic and NARMA-50 at ratio gamma."""
    _, yp  = make_periodic(T, seed)
    ua, ya = narma(order=50, T=T, seed=seed)
    h = (1 - gamma) * yp + gamma * ya
    h = (h - h.mean()) / (h.std() + 1e-12)
    u_mix = (1 - gamma) * np.roll(yp, 1) + gamma * ua
    u_mix = (u_mix - u_mix.mean()) / (u_mix.std() + 1e-12)
    return u_mix, h


def add_noise(u: np.ndarray, snr_db: float, seed: int = 0):
    """Add AWGN to u at given SNR in dB."""
    np.random.seed(seed + 300)
    sig_pow  = np.var(u)
    noise_pw = sig_pow / 10 ** (snr_db / 10)
    return u + np.random.randn(len(u)) * np.sqrt(noise_pw)


def run_kr(u, y, rho_f, rho_s, lam, alpha_s, M, seed):
    m = KRReservoir(N=N, M=M, rho_f=rho_f, rho_s=rho_s,
                    alpha_s=alpha_s, lam=lam, seed=seed, washout=WASHOUT)
    m.fit(u, y)
    n, _, _ = m.evaluate(u, y, TEST_SZ)
    return n


def run_esn(u, y, rho, sigma, lam, seed):
    m = StandardESN(N=N, rho=rho, sigma=sigma, lam=lam,
                    seed=seed, washout=WASHOUT)
    m.fit(u, y)
    n, _, _ = m.evaluate(u, y, TEST_SZ)
    return n


def gs_simple(u_h, y_h, method, M, alpha_s):
    """Quick grid search on half-data."""
    half = len(u_h) // 2; best = (np.inf, None)
    for rho, lam in itertools.product(RHO_GRID, LAM_GRID):
        if method == 'std':
            n = run_esn(u_h[:half], y_h[:half], rho, 0.1, lam, VAL_SEED)
            p = (rho, 0.1, lam)
        else:
            n = run_kr(u_h[:half], y_h[:half], rho, rho, lam, alpha_s, M, VAL_SEED)
            p = (rho, rho, lam)
        if n < best[0]: best = (n, p)
    return best[1]


# ════════════════════════════════════════════════════════════════
print("=" * 60)
print("EXPERIMENT 1: HYBRID SIGNAL CROSSOVER")
print("=" * 60)

T_HYB   = 3000
GAMMAS  = [0.00, 0.15, 0.30, 0.45, 0.60, 0.75, 0.90, 1.00]
esn_h   = []; kr_h = []; esn_sh = []; kr_sh = []

print(f"{'gamma':>6}  {'ESN':>12}  {'K-R':>12}  {'Impv%':>8}  Winner")
print("-" * 52)

for gamma in GAMMAS:
    u_hp, y_hp = make_hybrid(gamma, T_HYB, seed=VAL_SEED)
    sp = gs_simple(u_hp, y_hp, 'std', M_HYB, AS_HYB)
    kp = gs_simple(u_hp, y_hp, 'kr',  M_HYB, AS_HYB)

    e_s = []; k_s = []
    for seed in range(N_SEEDS):
        uv, yv = make_hybrid(gamma, T_HYB, seed=seed)
        e_s.append(run_esn(uv, yv, sp[0], sp[1], sp[2], seed))
        k_s.append(run_kr( uv, yv, kp[0], kp[1], kp[2], AS_HYB, M_HYB, seed))

    em, es = np.mean(e_s), np.std(e_s)
    km, ks = np.mean(k_s), np.std(k_s)
    esn_h.append(em); kr_h.append(km); esn_sh.append(es); kr_sh.append(ks)
    impv = (em - km) / em * 100
    winner = "K-R" if km < em else "ESN"
    print(f"{gamma:>6.2f}  {em:>8.4f}±{es:.3f}  {km:>8.4f}±{ks:.3f}  "
          f"{impv:>+7.1f}%  {winner}")

# ════════════════════════════════════════════════════════════════
print("\n\n" + "=" * 60)
print("EXPERIMENT 2: NOISE ROBUSTNESS (NARMA-50)")
print("=" * 60)

M_NR   = 50
AS_NR  = 0.05
T_NR   = 3000
SNR_DB = [30, 20, 15, 10, 5, 0, -5]

u_c, y_c = narma(M_NR, T_NR, seed=VAL_SEED)
sp_nr = gs_simple(u_c, y_c, 'std', M_NR, AS_NR)
kp_nr = gs_simple(u_c, y_c, 'kr',  M_NR, AS_NR)

# Clean baseline
eb_sc = []; kb_sc = []
for seed in range(N_SEEDS):
    uv, yv = narma(M_NR, T_NR, seed=seed)
    eb_sc.append(run_esn(uv, yv, sp_nr[0], sp_nr[1], sp_nr[2], seed))
    kb_sc.append(run_kr( uv, yv, kp_nr[0], kp_nr[1], kp_nr[2], AS_NR, M_NR, seed))
eb = np.mean(eb_sc); kb = np.mean(kb_sc)
print(f"{'Clean':>7}  ESN={eb:.4f}±{np.std(eb_sc):.3f}  "
      f"K-R={kb:.4f}±{np.std(kb_sc):.3f}  (baseline)")

esn_nr=[]; kr_nr=[]; esn_snr=[]; kr_snr=[]
for snr in SNR_DB:
    e_sc = []; k_sc = []
    for seed in range(N_SEEDS):
        uv, yv = narma(M_NR, T_NR, seed=seed)
        un = add_noise(uv, snr, seed)
        e_sc.append(run_esn(un, yv, sp_nr[0], sp_nr[1], sp_nr[2], seed))
        k_sc.append(run_kr( un, yv, kp_nr[0], kp_nr[1], kp_nr[2], AS_NR, M_NR, seed))
    em, es = np.mean(e_sc), np.std(e_sc)
    km, ks = np.mean(k_sc), np.std(k_sc)
    esn_nr.append(em); kr_nr.append(km)
    esn_snr.append(es); kr_snr.append(ks)
    de = (em - eb) / eb * 100; dk = (km - kb) / kb * 100
    print(f"{snr:>7}dB  ESN={em:.4f}±{es:.3f}  K-R={km:.4f}±{ks:.3f}  "
          f"ESN:{de:+.1f}%  K-R:{dk:+.1f}%  "
          f"{'ESN more robust' if de < dk else 'K-R more robust'}")


# ════════════════════════════════════════════════════════════════
# ── Figures ──────────────────────────────────────────────────
fig, axes = plt.subplots(1, 2, figsize=(11, 4))

# Figure H: Crossover
ax = axes[0]
ga = np.array(GAMMAS)
ea = np.array(esn_h); ka = np.array(kr_h)
es_a = np.array(esn_sh); ks_a = np.array(kr_sh)
ax.fill_between(ga, ea - es_a, ea + es_a, alpha=0.15, color='#4C72B0')
ax.fill_between(ga, ka - ks_a, ka + ks_a, alpha=0.15, color='#2CA02C')
ax.plot(ga, ea, 'o-', color='#4C72B0', lw=2, ms=6, label='Tuned ESN')
ax.plot(ga, ka, 's-', color='#2CA02C', lw=2, ms=6, label='K-R (proposed)')
for g, e, k in zip([0.15, 0.45, 1.0],
                   [esn_h[1], esn_h[3], esn_h[-1]],
                   [kr_h[1],  kr_h[3],  kr_h[-1]]):
    ax.annotate(f'+{(e-k)/e*100:.0f}%', xy=(g, (e+k)/2),
                fontsize=8, color='#1a6e1a', fontweight='bold',
                xytext=(5, 0), textcoords='offset points')
ax.set_xlabel('Aperiodic fraction γ  (0=periodic, 1=NARMA-50)', fontsize=10)
ax.set_ylabel('NRMSE', fontsize=10)
ax.set_title('Hybrid Signal Crossover\nK-R advantage grows continuously with γ',
             fontsize=10, fontweight='bold')
ax.legend(fontsize=9); ax.grid(alpha=0.3)
ax.tick_params(labelsize=9)

# Figure I: Noise robustness
ax2 = axes[1]
snr_x = np.array(SNR_DB)
ea_nr = np.array(esn_nr); ka_nr = np.array(kr_nr)
es_nr = np.array(esn_snr); ks_nr = np.array(kr_snr)
ax2.axhline(eb, color='#4C72B0', lw=1.2, ls=':', alpha=0.8, label='ESN (clean)')
ax2.axhline(kb, color='#2CA02C', lw=1.2, ls=':', alpha=0.8, label='K-R (clean)')
ax2.fill_between(snr_x, ea_nr - es_nr, ea_nr + es_nr, alpha=0.15, color='#4C72B0')
ax2.fill_between(snr_x, ka_nr - ks_nr, ka_nr + ks_nr, alpha=0.15, color='#2CA02C')
ax2.plot(snr_x, ea_nr, 'o-', color='#4C72B0', lw=2, ms=6, label='ESN (noisy)')
ax2.plot(snr_x, ka_nr, 's-', color='#2CA02C', lw=2, ms=6, label='K-R (noisy)')
ax2.axvline(10, color='#888888', lw=1.2, ls='--', alpha=0.7)
ax2.text(10.5, max(ea_nr) * 0.95, 'SNR=10dB\nthreshold',
         fontsize=8, color='#555555')
ax2.set_xlabel('Input SNR (dB)', fontsize=10)
ax2.set_ylabel('NRMSE', fontsize=10)
ax2.set_title('Noise Robustness (NARMA-50)\nK-R degrades faster below 10 dB',
             fontsize=10, fontweight='bold')
ax2.invert_xaxis(); ax2.legend(fontsize=9); ax2.grid(alpha=0.3)
ax2.tick_params(labelsize=9)

plt.tight_layout()
plt.savefig('figH_crossover_noise.png', dpi=180, bbox_inches='tight',
            facecolor='white')
plt.close()
print("\nFigure saved: figH_crossover_noise.png")

# Save results
with open('results_advanced.pkl', 'wb') as f:
    pickle.dump(dict(
        gammas=GAMMAS, esn_h=esn_h, kr_h=kr_h,
        esn_sh=esn_sh, kr_sh=kr_sh,
        snr_db=SNR_DB, esn_nr=esn_nr, kr_nr=kr_nr,
        esn_base=eb, kr_base=kb
    ), f)
print("Results saved: results_advanced.pkl")
