"""Subject/session-only EEGNet evaluation for TCH ASR20 + ICA data."""
from __future__ import annotations

import argparse
import json
from pathlib import Path

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

from event_stft.models import EEGNet
from event_stft.support_csp import normalize_task
from event_stft.training import resolve_device, seed_everything, stable_seed


ROOT = Path(__file__).resolve().parent
DEFAULT_CACHE = ROOT.parent / "TCH_Sequential_8ch" / "data_cache_asr20_ica"
SESSIONS = {"s1": ("s1_1", "s1_2", "s1_3", "s1_4", "s1_5"), "s2": ("s2_1", "s2_2", "s2_3", "s2_4")}


def augment(x: torch.Tensor, generator: torch.Generator) -> torch.Tensor:
    result = x.clone(); batch, channels, samples = result.shape
    result *= torch.empty((batch, channels, 1), device=x.device).uniform_(.95, 1.05, generator=generator)
    result += torch.empty((batch, channels, 1), device=x.device).uniform_(-.05, .05, generator=generator)
    shifts = torch.randint(-10, 11, (batch,), generator=generator, device=x.device)
    padded = torch.nn.functional.pad(result, (10, 10), mode="reflect")
    return torch.stack([padded[i, :, 10-int(shift):10-int(shift)+samples] for i, shift in enumerate(shifts)])


@torch.no_grad()
def accuracy(model: nn.Module, x: np.ndarray, y: np.ndarray, device: torch.device) -> float:
    model.eval(); prediction = model(torch.as_tensor(x, dtype=torch.float32, device=device)).argmax(-1).cpu().numpy()
    return float((prediction == y).mean())


def train_fold(x: np.ndarray, y: np.ndarray, train_indices: np.ndarray, val_indices: np.ndarray, device: torch.device, seed: int) -> tuple[nn.Module, int, float]:
    seed_everything(seed)
    model = EEGNet(channels=8, samples=400, dropout=.5).to(device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-3)
    loss_fn = nn.CrossEntropyLoss(label_smoothing=.05)
    generator = torch.Generator(device=device.type).manual_seed(seed)
    train_x = torch.as_tensor(x[train_indices], dtype=torch.float32, device=device); train_y = torch.as_tensor(y[train_indices], dtype=torch.long, device=device)
    best_state, best_val, patience, stale = None, -np.inf, 20, 0
    for epoch in range(1, 151):
        model.train()
        for batch in torch.randperm(len(train_y), generator=generator, device=device).split(min(16, len(train_y))):
            optimizer.zero_grad(set_to_none=True); loss = loss_fn(model(augment(train_x[batch], generator)), train_y[batch]); loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); optimizer.step()
        value = accuracy(model, x[val_indices], y[val_indices], device)
        if value > best_val:
            best_val, best_state, stale = value, {key:value.detach().cpu().clone() for key,value in model.state_dict().items()}, 0
        else:
            stale += 1
            if stale >= patience: break
    model.load_state_dict(best_state)
    return model, epoch, best_val


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--cache", type=Path, default=DEFAULT_CACHE)
    parser.add_argument("--output", type=Path, default=ROOT / "tch_asr20_ica_eegnet_self_runs")
    parser.add_argument("--seed", type=int, default=20260721)
    parser.add_argument("--device", default="cuda")
    args = parser.parse_args(); seed_everything(args.seed); device = resolve_device(args.device); args.output.mkdir(parents=True, exist_ok=True)
    fold_rows, session_rows = [], []
    for subject, sessions in SESSIONS.items():
        for order, session in enumerate(sessions, 1):
            epochs = np.load(args.cache / "sessions" / f"{session}_x.npy", allow_pickle=False)
            y = np.load(args.cache / "sessions" / f"{session}_y.npy", allow_pickle=False).astype(np.int64)
            x = normalize_task(epochs[:, :, 200:600])
            splitter = StratifiedKFold(5, shuffle=True, random_state=stable_seed(args.seed, subject, session, "outer"))
            scores = []
            for fold, (train_index, test_index) in enumerate(splitter.split(x, y), 1):
                inner_train, val_index = train_test_split(train_index, test_size=.2, stratify=y[train_index], random_state=stable_seed(args.seed, subject, session, fold, "validation"))
                model, epochs_run, best_val = train_fold(x, y, inner_train, val_index, device, stable_seed(args.seed, subject, session, fold, "model"))
                score = accuracy(model, x[test_index], y[test_index], device); scores.append(score)
                fold_rows.append({"subject_id":subject,"session_id":session,"fold":fold,"test_accuracy":score,"validation_accuracy":best_val,"epochs_run":epochs_run,"train_trials":len(inner_train),"validation_trials":len(val_index),"test_trials":len(test_index)})
            session_rows.append({"subject_id":subject,"session_id":session,"session_order":order,"model":"EEGNet scratch on TCH ASR20+ICA","mean_accuracy":float(np.mean(scores)),"std_accuracy":float(np.std(scores,ddof=1)),"trials":len(y),"left":int((y==0).sum()),"right":int((y==1).sum())})
            print(f"[self EEGNet] {session}: {np.mean(scores):.3f}", flush=True)
    pd.DataFrame(fold_rows).to_csv(args.output / "fold_results.csv", index=False, encoding="utf-8-sig")
    pd.DataFrame(session_rows).to_csv(args.output / "session_results.csv", index=False, encoding="utf-8-sig")
    pd.DataFrame(session_rows).groupby("subject_id", as_index=False).agg(mean_accuracy=("mean_accuracy","mean"),std_across_sessions=("mean_accuracy","std"),sessions=("session_id","count")).to_csv(args.output / "subject_summary.csv", index=False, encoding="utf-8-sig")
    (args.output / "run_config.json").write_text(json.dumps({"data":"TCH ASR cutoff 20 + conservative ICA cache only", "model":"EEGNet random initialization per outer fold", "evaluation":"each session independently uses stratified outer 5-fold; outer-train has an inner 80/20 validation split for early stopping; no 34-subject source data, no checkpoint, no cross-session carry-over", "seed":args.seed},ensure_ascii=False,indent=2), encoding="utf-8")


if __name__ == "__main__": main()
