"""Leakage-controlled self-only evaluation on the TCH v2 cleaned cache.

Each recording session is an independent task.  No source-subject checkpoint,
other session, or outer-test trial is used for training, CSP fitting, or early
stopping.  The v2 cleaner is deliberately treated as a fixed signal transform;
this script only evaluates its resulting epochs.
"""
from __future__ import annotations

import argparse
import json
from pathlib import Path

import numpy as np
import pandas as pd
from sklearn.model_selection import StratifiedKFold, train_test_split

from event_stft.support_csp import CSPConfig, train_target_fold
from event_stft.training import TrainConfig, stable_seed, train_fold


ROOT = Path(__file__).resolve().parent
DEFAULT_CACHE = ROOT.parent / "TCH_Sequential_8ch" / "data_cache_v2_cleaned"
SESSIONS = ("s1_1", "s1_2", "s1_3", "s1_4", "s1_5", "s2_1", "s2_2", "s2_3", "s2_4")
MODEL_NAMES = {"eegnet_task": "EEGNet", "shallow_task": "ShallowConvNet"}


def make_split(y: np.ndarray, outer_train: np.ndarray, outer_test: np.ndarray, seed: int) -> dict[str, np.ndarray]:
    train, val = train_test_split(
        outer_train,
        test_size=0.20,
        stratify=y[outer_train],
        random_state=seed,
    )
    return {"train": np.asarray(train), "val": np.asarray(val), "test": np.asarray(outer_test)}


def run_session(session: str, x: np.ndarray, y: np.ndarray, output: Path, seed: int, device: str) -> list[dict]:
    rows: list[dict] = []
    splitter = StratifiedKFold(n_splits=5, shuffle=True, random_state=stable_seed(seed, session, "outer"))
    deep_config = TrainConfig(
        device=device,
        seed=seed,
        max_epochs=150,
        patience=20,
        batch_size_all=16,
        batch_size_fewshot=8,
        learning_rate=3e-4,
        weight_decay=1e-3,
        label_smoothing=0.05,
        amplitude_low=0.95,
        amplitude_high=1.05,
        max_shift_samples=10,
    )
    csp_config = CSPConfig(
        device=device,
        seed=seed,
        target_epochs=150,
        target_patience=20,
        batch_size=16,
        fewshot_batch_size=8,
        finetune_lr=3e-4,
        weight_decay=1e-3,
        label_smoothing=0.05,
        bootstrap_references=1,
        crossfit_train_queries=True,
    )
    for fold, (outer_train, outer_test) in enumerate(splitter.split(x, y), 1):
        split = make_split(y, outer_train, outer_test, stable_seed(seed, session, fold, "validation"))
        for model_name, display_name in MODEL_NAMES.items():
            result = train_fold(
                x=x,
                y=y,
                split=split,
                model_name=model_name,
                representation="both",
                train_seed=stable_seed(seed, session, fold, model_name),
                config=deep_config,
                max_epochs=150,
                patience=20,
                evaluate_test=True,
                checkpoint_path=None,
            )
            rows.append({**result, "session_id": session, "fold": fold, "model": display_name})
            print(f"[{session} fold {fold}] {display_name}: {result['accuracy']:.3f}", flush=True)

        result = train_target_fold(
            x=x,
            y=y,
            split=split,
            mode="Support-CSP",
            seed=stable_seed(seed, session, fold, "support_csp"),
            config=csp_config,
            initial_state=None,
            evaluate_test=True,
            checkpoint=None,
        )
        rows.append({**result, "session_id": session, "fold": fold, "model": "Support-CSP"})
        print(f"[{session} fold {fold}] Support-CSP: {result['accuracy']:.3f}", flush=True)
    return rows


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--cache", type=Path, default=DEFAULT_CACHE)
    parser.add_argument("--output", type=Path, default=ROOT / "tch_v2_clean_three_model_runs")
    parser.add_argument("--seed", type=int, default=20260722)
    parser.add_argument("--device", default="cuda")
    args = parser.parse_args()
    cleaned = args.cache / "cleaned_sessions"
    if not cleaned.exists():
        raise FileNotFoundError(f"Expected cleaned_sessions under {args.cache}")
    args.output.mkdir(parents=True, exist_ok=True)

    rows: list[dict] = []
    manifest: list[dict] = []
    for session in SESSIONS:
        x_path, y_path = cleaned / f"{session}_x.npy", cleaned / f"{session}_y.npy"
        if not x_path.exists() or not y_path.exists():
            raise FileNotFoundError(f"Missing cache files for {session}")
        x, y = np.load(x_path, allow_pickle=False).astype(np.float32), np.load(y_path, allow_pickle=False).astype(np.int64)
        if x.ndim != 3 or x.shape[1:] != (8, 600) or len(x) != len(y) or not np.isfinite(x).all():
            raise ValueError(f"Invalid cache for {session}: x={x.shape}, y={y.shape}")
        if min(np.bincount(y, minlength=2)) < 10:
            raise ValueError(f"Too few trials per class in {session}")
        manifest.append({"session_id": session, "trials": len(y), "left": int((y == 0).sum()), "right": int((y == 1).sum()), "shape": str(tuple(x.shape))})
        rows.extend(run_session(session, x, y, args.output, args.seed, args.device))
        pd.DataFrame(rows).to_csv(args.output / "fold_results.csv", index=False, encoding="utf-8-sig")

    folds = pd.DataFrame(rows)
    summary = folds.groupby(["session_id", "model"], as_index=False).agg(
        mean_accuracy=("accuracy", "mean"),
        std_accuracy=("accuracy", "std"),
        mean_balanced_accuracy=("balanced_accuracy", "mean"),
        mean_macro_f1=("macro_f1", "mean"),
        mean_best_epoch=("best_epoch", "mean"),
        parameters=("parameter_count", "first"),
    )
    group = summary.groupby("model", as_index=False).agg(
        session_mean_accuracy=("mean_accuracy", "mean"),
        session_std_accuracy=("mean_accuracy", "std"),
        mean_balanced_accuracy=("mean_balanced_accuracy", "mean"),
        mean_macro_f1=("mean_macro_f1", "mean"),
        sessions=("session_id", "count"),
    )
    pd.DataFrame(manifest).to_csv(args.output / "data_manifest.csv", index=False, encoding="utf-8-sig")
    summary.to_csv(args.output / "session_summary.csv", index=False, encoding="utf-8-sig")
    group.to_csv(args.output / "group_summary.csv", index=False, encoding="utf-8-sig")
    config = {
        "data": "TCH v2 cleaned sessions only; no source normal subjects, no cross-session carry-over",
        "cleaning": "despike (one-sample gradient >200 in raw units), continuous 4th-order zero-phase Butterworth 4-32 Hz, session channel z-score, resample 1000->200 Hz, event[-1,2] epoch, [-1,0] mean baseline subtraction",
        "model_input": "task [0,2] s (8,400); each trial/channel z-score is applied inside all three models",
        "evaluation": "each session independently: stratified outer 5-fold; outer-train split 80/20 for validation early stopping; outer test never selects epoch or CSP parameters",
        "support_csp": "three bands 8-13/13-20/20-30 Hz, shrinkage CSP fit only on inner-train, cross-fitted CSP features for inner-train, val/test transformed by inner-train CSP",
        "seed": args.seed,
    }
    (args.output / "run_config.json").write_text(json.dumps(config, ensure_ascii=False, indent=2), encoding="utf-8")


if __name__ == "__main__":
    main()
