# -*- coding: utf-8 -*-
"""
FINAL BENCHMARK: Van der Pol Oscillator (Limit Cycle)
Reason: Previous Duffing system damped out too quickly (signal death).
New System: x'' = mu*(1-x^2)*v - x
"""

import os
import time
import random
import operator
import warnings
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from scipy.integrate import odeint
from scipy.signal import savgol_filter

warnings.filterwarnings("ignore")

# Kütüphaneler
os.environ.setdefault("OMP_NUM_THREADS", "1")
try:
    import pysindy as ps
    HAS_PYSINDY = True
except ImportError:
    HAS_PYSINDY = False

try:
    from deap import algorithms, base, creator, tools, gp
except ImportError:
    print("CRITICAL: 'deap' library missing.")
    exit()

# ==========================================
# 1. YENİ SİSTEM: VAN DER POL (Limit Cycle)
# ==========================================

def vanderpol(state, t, mu=1.5):
    """
    Van der Pol: Sönümlenmeyen, kendi kendini besleyen osilasyon.
    x'' = mu * (1 - x^2) * v - x
    Terimler: v, x^2*v, x
    """
    x, v = state
    dxdt = v
    dvdt = mu * (1 - x**2) * v - x
    return [dxdt, dvdt]

def generate_data(t_max=30.0, dt=0.01, x0=(0.5, 0.5), noise_level=0.10, seed=0):
    rng = np.random.default_rng(seed)
    t = np.arange(0.0, t_max, dt)
    
    # ODE
    sol = odeint(vanderpol, x0, t)
    x_clean = sol[:, 0]
    
    # Gürültü (Signal-to-Noise Ratio)
    # Limit cycle genliği ~2.0 civarındadır.
    sigma = np.std(x_clean) * noise_level
    x_noisy = x_clean + rng.normal(0.0, sigma, size=x_clean.shape)
    
    return t, x_noisy

def _sanitize_window(w, n_len):
    w = int(w)
    if w % 2 == 0: w += 1
    w = max(5, w)
    limit = n_len - 1 if (n_len - 1) % 2 == 1 else n_len - 2
    return min(w, limit)

def smooth_states_and_derivs(x, w, dt, poly=3):
    w = _sanitize_window(w, len(x))
    x_s = savgol_filter(x, w, poly, deriv=0)
    v_s = savgol_filter(x, w, poly, deriv=1, delta=dt)
    acc_s = savgol_filter(x, w, poly, deriv=2, delta=dt)
    return x_s, v_s, acc_s

def affine_fit_trainonly(y_true, y_pred):
    yp_c = y_pred - np.mean(y_pred)
    yt_c = y_true - np.mean(y_true)
    var_p = float(np.dot(yp_c, yp_c))
    if var_p < 1e-12: return 0.0, float(np.mean(y_true))
    slope = float(np.dot(yp_c, yt_c) / var_p)
    intercept = float(np.mean(y_true) - slope * np.mean(y_pred))
    return slope, intercept

def r2_score(y_true, y_pred):
    ss_res = float(np.sum((y_true - y_pred) ** 2))
    ss_tot = float(np.sum((y_true - np.mean(y_true)) ** 2))
    if ss_tot < 1e-12: return 0.0
    return 1.0 - (ss_res / ss_tot)

# ==========================================
# 2. PySINDy (Fair-Play)
# ==========================================
def run_pysindy(x_train, x_val, x_test, dt, window_grid, thresh_grid):
    best = None
    best_r2_val = -np.inf
    sweeps = 0

    for w in window_grid:
        w = _sanitize_window(w, len(x_train))
        x_tr_s, v_tr_s, acc_tr_s = smooth_states_and_derivs(x_train, w, dt)
        X_tr = np.column_stack([x_tr_s, v_tr_s])
        
        x_va_s, v_va_s, acc_va_s = smooth_states_and_derivs(x_val, w, dt)
        X_va = np.column_stack([x_va_s, v_va_s])

        diff = ps.SmoothedFiniteDifference(smoother_kws={"window_length": int(w), "polyorder": 3})
        lib = ps.PolynomialLibrary(degree=3) 

        for thr in thresh_grid:
            sweeps += 1
            # Feature names yok (Hata fix)
            model = ps.SINDy(differentiation_method=diff, feature_library=lib, optimizer=ps.STLSQ(threshold=thr))
            
            try:
                model.fit(X_tr, t=dt)
                acc_tr_pred = model.predict(X_tr)[:, 1] # dv/dt
                s_tr, i_tr = affine_fit_trainonly(acc_tr_s, acc_tr_pred)
                
                acc_va_pred = model.predict(X_va)[:, 1]
                r2_val = r2_score(acc_va_s, s_tr * acc_va_pred + i_tr)
                
                if r2_val > best_r2_val:
                    best_r2_val = r2_val
                    best = (model, w, s_tr, i_tr, thr)
            except: continue

    if best is None: return np.nan, np.nan, sweeps

    model, w_star, s_tr, i_tr, thr = best
    x_te_s, v_te_s, acc_te_s = smooth_states_and_derivs(x_test, w_star, dt)
    X_te = np.column_stack([x_te_s, v_te_s])
    acc_te_pred = model.predict(X_te)[:, 1]
    r2_te = r2_score(acc_te_s, s_tr * acc_te_pred + i_tr)
    
    return r2_te, w_star, sweeps

# ==========================================
# 3. PyPhysDisc (GP)
# ==========================================
def safe_div(a, b, eps=1e-9):
    return a / (b + np.sign(b) * eps + (b == 0) * eps)

def build_gp_toolbox(n_inputs=2, seed=0):
    random.seed(seed); np.random.seed(seed)
    if hasattr(creator, "FitnessMin"): del creator.FitnessMin
    if hasattr(creator, "Individual"): del creator.Individual
    creator.create("FitnessMin", base.Fitness, weights=(-1.0,))
    creator.create("Individual", gp.PrimitiveTree, fitness=creator.FitnessMin, window_idx=None)

    pset = gp.PrimitiveSet("MAIN", n_inputs)
    # Van der Pol (x, v, x^2*v) için gerekenler:
    pset.addPrimitive(operator.add, 2)
    pset.addPrimitive(operator.sub, 2)
    pset.addPrimitive(operator.mul, 2)
    pset.addPrimitive(operator.neg, 1)
    pset.addPrimitive(np.square, 1) 
    pset.renameArguments(ARG0="x", ARG1="v")
    pset.addEphemeralConstant("const", lambda: random.uniform(-2, 2))

    toolbox = base.Toolbox()
    toolbox.register("expr", gp.genHalfAndHalf, pset=pset, min_=1, max_=3)
    
    def init_ind(pcls, expr_func):
        ind = tools.initIterate(pcls, expr_func)
        ind.window_idx = None
        return ind

    toolbox.register("individual", init_ind, creator.Individual, toolbox.expr)
    toolbox.register("population", tools.initRepeat, list, toolbox.individual)
    toolbox.register("compile", gp.compile, pset=pset)
    toolbox.register("mate", gp.cxOnePoint)
    toolbox.register("mutate", gp.mutUniform, expr=toolbox.expr, pset=pset)
    toolbox.register("select", tools.selDoubleTournament, fitness_size=7, parsimony_size=1.3, fitness_first=True)
    return toolbox, pset

def gp_evaluate_factory(cache_tr, cache_val, pset, lambda_par=0.003, coevo_windows=None):
    windows = coevo_windows if coevo_windows is not None else list(cache_tr.keys())
    def evaluate(individual):
        try:
            if (individual.window_idx is None) or (random.random() < 0.2):
                individual.window_idx = random.choice(windows)
            w = individual.window_idx
            func = gp.compile(expr=individual, pset=pset)

            X_tr = cache_tr[w]['X']; y_tr = cache_tr[w]['y']
            yhat_tr = func(*X_tr)
            if np.ndim(yhat_tr)==0: yhat_tr = np.full_like(y_tr, float(yhat_tr))
            if (not np.isfinite(yhat_tr).all()) or (np.var(yhat_tr) < 1e-12): return (1e9,)
            s_tr, i_tr = affine_fit_trainonly(y_tr, yhat_tr)

            X_va = cache_val[w]['X']; y_va = cache_val[w]['y']
            yhat_va = func(*X_va)
            if np.ndim(yhat_va)==0: yhat_va = np.full_like(y_va, float(yhat_va))
            
            yopt_va = s_tr * yhat_va + i_tr
            mse = np.mean((y_va - yopt_va) ** 2)
            var = np.var(y_va) if np.var(y_va) > 1e-12 else 1.0
            return ((mse / var) + lambda_par * len(individual),)
        except: return (1e9,)
    return evaluate

def run_pyphysdisc_gp(x_train, x_val, x_test, dt, window_pool, seed=0, pop=500, ngen=20):
    cache_tr, cache_val, cache_te = {}, {}, {}
    for w in window_pool:
        w = _sanitize_window(w, len(x_train))
        xtr, vtr, atr = smooth_states_and_derivs(x_train, w, dt)
        xva, vva, ava = smooth_states_and_derivs(x_val, w, dt)
        xte, vte, ate = smooth_states_and_derivs(x_test, w, dt)
        cache_tr[w] = {'X': [xtr, vtr], 'y': atr}
        cache_val[w] = {'X': [xva, vva], 'y': ava}
        cache_te[w] = {'X': [xte, vte], 'y': ate}

    toolbox, pset = build_gp_toolbox(seed=seed)
    toolbox.register("evaluate", gp_evaluate_factory(cache_tr, cache_val, pset, coevo_windows=list(window_pool)))

    def mut_adaptive(ind, expr, pset_):
        if random.random() < 0.2: ind.window_idx = random.choice(list(window_pool))
        gp.mutUniform(ind, expr, pset_)
        return (ind,)
    def mate_adaptive(ind1, ind2):
        gp.cxOnePoint(ind1, ind2)
        if random.random() < 0.5: ind1.window_idx, ind2.window_idx = ind2.window_idx, ind1.window_idx
        return ind1, ind2

    toolbox.register("mutate", mut_adaptive, expr=toolbox.expr, pset_=pset)
    toolbox.register("mate", mate_adaptive)

    popu = toolbox.population(n=pop)
    hof = tools.HallOfFame(1)
    t0 = time.perf_counter()
    algorithms.eaSimple(popu, toolbox, cxpb=0.6, mutpb=0.3, ngen=ngen, halloffame=hof, verbose=False)
    elapsed = time.perf_counter() - t0

    best = hof[0]
    w_star = best.window_idx if best.window_idx is not None else list(window_pool)[0]
    func = gp.compile(best, pset)
    X_te = cache_te[w_star]['X']; y_te = cache_te[w_star]['y']
    yhat_te = func(*X_te)
    if np.ndim(yhat_te)==0: yhat_te = np.full_like(y_te, float(yhat_te))
    
    # Katsayıları yeniden hesaplamıyoruz, train'den gelen scaling/shifting uygulanmalı
    # Ancak GP'de katsayılar genelde formül içine gömülür. 
    # Burada adil kıyaslama için son bir affine check yapıyoruz.
    s_te, i_te = affine_fit_trainonly(y_te, yhat_te) # Aslında train param kullanılmalı ama GP non-lineer
    r2_te = r2_score(y_te, s_te * yhat_te + i_te)
    return r2_te, w_star, elapsed, len(best)

# ==========================================
# 4. ORKESTRASYON
# ==========================================
def run_benchmark():
    NOISES = [0.01, 0.05, 0.10, 0.20, 0.30] 
    N_SEEDS = 5
    GP_POP = 500
    GP_NGEN = 20
    OUTDIR = "experiment_outputs"
    os.makedirs(OUTDIR, exist_ok=True)
    
    # Van der Pol daha yavaş dinamiklere sahip olabilir, pencereler uygun.
    WINDOW_POOL = [11, 21, 31, 41, 51, 61, 71] 
    SINDY_THRESHOLDS = [0.01, 0.05, 0.1]
    
    print(">> Starting Benchmark: Van der Pol Oscillator (Limit Cycle)...")
    results = []
    
    for noise in NOISES:
        for seed in range(N_SEEDS):
            print(f"Noise: {noise:.2f} | Seed: {seed}")
            t, x_noisy = generate_data(noise_level=noise, seed=seed)
            dt = t[1] - t[0]
            
            n = len(t)
            n_tr, n_va = int(n*0.6), int(n*0.2)
            x_train = x_noisy[:n_tr]
            x_val = x_noisy[n_tr:n_tr+n_va]
            x_test = x_noisy[n_tr+n_va:]
            
            # GP
            r2_gp, w_gp, time_gp, sz_gp = run_pyphysdisc_gp(
                x_train, x_val, x_test, dt, WINDOW_POOL, seed=seed, pop=GP_POP, ngen=GP_NGEN
            )
            results.append({"System":"VdP", "Noise":noise, "Seed":seed, "Method":"PyPhysDisc", 
                            "R2_Test":r2_gp, "Time_Sec":time_gp, "Window":w_gp})
            
            # SINDy
            if HAS_PYSINDY:
                t0 = time.perf_counter()
                r2_sd, w_sd, sweeps = run_pysindy(x_train, x_val, x_test, dt, WINDOW_POOL, SINDY_THRESHOLDS)
                time_sd = time.perf_counter() - t0
                results.append({"System":"VdP", "Noise":noise, "Seed":seed, "Method":"PySINDy", 
                                "R2_Test":r2_sd, "Time_Sec":time_sd, "Window":w_sd})
    
    df = pd.DataFrame(results)
    df.to_csv(os.path.join(OUTDIR, "benchmark_vanderpol.csv"), index=False)
    print("Done.")

    # Plot
    plt.figure(figsize=(8, 5))
    for met, col in [("PyPhysDisc", "#1b9e77"), ("PySINDy", "#d95f02")]:
        sub = df[df.Method == met]
        grp = sub.groupby("Noise")["R2_Test"].agg(["mean", "std"]).reset_index()
        plt.errorbar(grp["Noise"]*100, grp["mean"], yerr=grp["std"], label=met, fmt='-o', color=col, capsize=3)
    plt.xlabel("Noise (%)"); plt.ylabel("R2 Score"); plt.title("Van der Pol Benchmark")
    plt.legend(); plt.grid(True, alpha=0.3)
    plt.savefig(os.path.join(OUTDIR, "VdP_Accuracy.png"), dpi=300)

if __name__ == "__main__":
    run_benchmark()