"""
run_main_experiments.py
=======================
Reproduce all main benchmark results from the paper.

Experiments:
    1. Main results: NARMA-10/30/50/100, Mackey-Glass, Santa Fe
    2. Ablation study (NARMA-10)
    3. Reservoir size scaling (NARMA-10)

Results are printed to console and saved to results_main.pkl.

Run:
    python run_main_experiments.py

Requirements:
    numpy, scipy
    benchmarks.py, kr_reservoir.py, esn_baseline.py  (same directory)
"""

import numpy as np
import pickle
from scipy import stats

from benchmarks import narma, mackey_glass, santa_fe
from kr_reservoir import KRReservoir
from esn_baseline import StandardESN, grid_search, evaluate_10seed

# ── Configuration ─────────────────────────────────────────────
N        = 200
N_SEEDS  = 10
TEST_SZ  = 500
WASHOUT  = 200
VAL_SEED = 99

TASKS = [
    ('NARMA-10',     'narma',  {'order': 10,  'T': 3000}),
    ('NARMA-30',     'narma',  {'order': 30,  'T': 3000}),
    ('NARMA-50',     'narma',  {'order': 50,  'T': 3000}),
    ('NARMA-100',    'narma',  {'order': 100, 'T': 4000}),
    ('Mackey-Glass', 'mg',     {'T': 3000}),
    ('Santa Fe',     'sf',     {'T': 3000}),
]

M_MAP = {
    'NARMA-10': 10, 'NARMA-30': 30,
    'NARMA-50': 50, 'NARMA-100': 100,
    'Mackey-Glass': 17, 'Santa Fe': 17,
}


def get_data(task_type, kwargs, seed=0):
    if task_type == 'narma':
        return narma(seed=seed, **kwargs)
    elif task_type == 'mg':
        return mackey_glass(seed=seed, **kwargs)
    elif task_type == 'sf':
        return santa_fe(seed=seed, **kwargs)


def sig_str(p):
    if p < 0.001: return '***'
    elif p < 0.01: return '**'
    elif p < 0.05: return '*'
    return 'ns'


# ════════════════════════════════════════════════════════════════
print("=" * 65)
print("K-R RESERVOIR ARCHITECTURE — MAIN RESULTS")
print("=" * 65)

all_results = {}

for name, task_type, task_kwargs in TASKS:
    M = M_MAP[name]
    print(f"\n--- {name} (M={M}) ---")

    # Get data using val_seed for grid search
    u_val, y_val = get_data(task_type, task_kwargs, seed=VAL_SEED)

    # Grid search
    std_params = grid_search(u_val, y_val, method='std', N=N, M=M,
                             washout=WASHOUT, test_size=200)
    kr_params  = grid_search(u_val, y_val, method='kr',  N=N, M=M,
                             washout=WASHOUT, test_size=200)
    print(f"  ESN best: rho={std_params['rho']:.2f} "
          f"sigma={std_params['sigma']:.2f} lam={std_params['lam']:.0e}")
    print(f"  K-R best: rho_f={kr_params['rho_f']:.2f} "
          f"rho_s={kr_params['rho_s']:.2f} lam={kr_params['lam']:.0e}")

    # 10-seed evaluation
    esn_scores = []
    kr_scores  = []
    for seed in range(N_SEEDS):
        u_ev, y_ev = get_data(task_type, task_kwargs, seed=seed)

        m_std = StandardESN(N=N, rho=std_params['rho'],
                            sigma=std_params['sigma'],
                            lam=std_params['lam'],
                            seed=seed, washout=WASHOUT)
        m_std.fit(u_ev, y_ev)
        ns, _, _ = m_std.evaluate(u_ev, y_ev, TEST_SZ)

        m_kr = KRReservoir(N=N, M=M, rho_f=kr_params['rho_f'],
                           rho_s=kr_params['rho_s'],
                           lam=kr_params['lam'],
                           seed=seed, washout=WASHOUT)
        m_kr.fit(u_ev, y_ev)
        nk, _, _ = m_kr.evaluate(u_ev, y_ev, TEST_SZ)

        esn_scores.append(ns)
        kr_scores.append(nk)

    em, es = np.mean(esn_scores), np.std(esn_scores)
    km, ks = np.mean(kr_scores),  np.std(kr_scores)
    impv    = (em - km) / em * 100
    _, p    = stats.ttest_ind(esn_scores, kr_scores)

    print(f"  ESN: {em:.4f} ± {es:.4f}")
    print(f"  K-R: {km:.4f} ± {ks:.4f}  Δ={impv:+.1f}%  {sig_str(p)}")

    all_results[name] = dict(
        esn_mean=em, esn_std=es, kr_mean=km, kr_std=ks,
        improvement=impv, p=p, sig=sig_str(p),
        esn_scores=esn_scores, kr_scores=kr_scores
    )


# ════════════════════════════════════════════════════════════════
print("\n\n" + "=" * 65)
print("ABLATION STUDY (NARMA-10, N=200)")
print("=" * 65)

u_bl, y_bl = narma(order=10, T=3000, seed=VAL_SEED)
# Untuned baseline
scores_ref = []
for seed in range(N_SEEDS):
    u_ev, y_ev = narma(order=10, T=3000, seed=seed)
    m = StandardESN(N=N, rho=0.9, sigma=1.0, lam=1e-4, seed=seed)
    m.fit(u_ev, y_ev); n, _, _ = m.evaluate(u_ev, y_ev, TEST_SZ)
    scores_ref.append(n)
print(f"  Standard ESN (untuned):  {np.mean(scores_ref):.4f} ± {np.std(scores_ref):.4f}")

ablation_configs = [
    ('+ Adaptive-K only',    True,  False, False),
    ('+ Leaky state only',   False, True,  False),
    ('+ Input delays only',  False, False, True),
    ('K-R Full (tuned)',     True,  True,  True),
]

from kr_reservoir import adaptive_k, make_reservoir, make_input_weights, adaptive_delays

M_AB = 10
DELAYS_AB = adaptive_delays(M_AB, d=10)
AS_AB = min(3.0 / M_AB, 1.0)
KF = adaptive_k(N); KS = adaptive_k(N, K0=1.0)

for label, use_k, use_slow, use_delays in ablation_configs:
    ab_scores = []
    for seed in range(N_SEEDS):
        u_ev, y_ev = narma(order=10, T=3000, seed=seed)

        if label == 'K-R Full (tuned)':
            # Use tuned K-R
            std_p = grid_search(u_bl, y_bl, 'std', N, M_AB, WASHOUT, 200)
            kr_p  = grid_search(u_bl, y_bl, 'kr',  N, M_AB, WASHOUT, 200)
            m = KRReservoir(N=N, M=M_AB, rho_f=kr_p['rho_f'],
                            rho_s=kr_p['rho_s'], lam=kr_p['lam'],
                            seed=seed, washout=WASHOUT)
            m.fit(u_ev, y_ev); n, _, _ = m.evaluate(u_ev, y_ev, TEST_SZ)
        else:
            # Manual ablation: partial K-R
            np.random.seed(seed)
            W  = make_reservoir(N, 0.9, seed)
            Wi = (make_input_weights(N, KF, seed) if use_k
                  else make_input_weights(N, 1.0, seed))

            if use_delays:
                md = max(DELAYS_AB) + 1
            x  = np.zeros(N); xs = np.zeros(N)
            db = np.zeros(max(DELAYS_AB) + 1) if use_delays else None
            Xs = []
            for t in range(len(u_ev)):
                if use_slow:
                    x = (1 - AS_AB) * x + AS_AB * np.tanh(W @ x + Wi * u_ev[t])
                else:
                    x = np.tanh(W @ x + Wi * u_ev[t])
                if use_delays:
                    db = np.roll(db, 1); db[0] = u_ev[t]
                    Xs.append(np.concatenate([x, db[DELAYS_AB]]))
                else:
                    Xs.append(x.copy())
            X = np.array(Xs)
            mu = X[WASHOUT:].mean(0); sg = X[WASHOUT:].std(0) + 1e-8
            X = (X - mu) / sg
            Xtr = X[WASHOUT:-TEST_SZ]; ytr = y_ev[WASHOUT:-TEST_SZ]
            Xte = X[-TEST_SZ:];       yte = y_ev[-TEST_SZ:]
            A   = Xtr.T @ Xtr + 1e-4 * np.eye(Xtr.shape[1])
            Wo  = np.linalg.solve(A, Xtr.T @ ytr)
            yhat = Xte @ Wo
            n = float(np.sqrt(np.mean((yte - yhat) ** 2)) / (np.std(yte) + 1e-12))
        ab_scores.append(n)
    print(f"  {label:<28s}  {np.mean(ab_scores):.4f} ± {np.std(ab_scores):.4f}")


# ════════════════════════════════════════════════════════════════
print("\n\n" + "=" * 65)
print("RESERVOIR SIZE SCALING (NARMA-10)")
print("=" * 65)

for N_test in [50, 100, 200, 500]:
    KF_n = adaptive_k(N_test)
    u_v, y_v = narma(order=10, T=3000, seed=VAL_SEED)
    sp = grid_search(u_v, y_v, 'std', N_test, 10, WASHOUT, 200)
    kp = grid_search(u_v, y_v, 'kr',  N_test, 10, WASHOUT, 200)

    e_sc = []; k_sc = []
    for seed in range(N_SEEDS):
        u_ev, y_ev = narma(order=10, T=3000, seed=seed)
        ms = StandardESN(N=N_test, rho=sp['rho'], sigma=sp['sigma'],
                         lam=sp['lam'], seed=seed, washout=WASHOUT)
        ms.fit(u_ev, y_ev); ns, _, _ = ms.evaluate(u_ev, y_ev, TEST_SZ)
        mk = KRReservoir(N=N_test, M=10, rho_f=kp['rho_f'],
                         rho_s=kp['rho_s'], lam=kp['lam'],
                         seed=seed, washout=WASHOUT)
        mk.fit(u_ev, y_ev); nk, _, _ = mk.evaluate(u_ev, y_ev, TEST_SZ)
        e_sc.append(ns); k_sc.append(nk)

    em, es = np.mean(e_sc), np.std(e_sc)
    km, ks = np.mean(k_sc), np.std(k_sc)
    print(f"  N={N_test:4d}  ESN={em:.4f}±{es:.3f}  "
          f"K-R={km:.4f}±{ks:.3f}  Δ={(em-km)/em*100:+.1f}%")


# ── Save all results ──────────────────────────────────────────
with open('results_main.pkl', 'wb') as f:
    pickle.dump(all_results, f)
print("\n\nAll results saved to results_main.pkl")
