"""Sequential cross-device evaluation for the Taichung Veterans General Hospital data."""
from __future__ import annotations

import argparse
import copy
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

from event_stft.models import EEGNet, ShallowConvNet, count_parameters
from event_stft.support_csp import CSPConfig, CSPReference, SupportCSPShallowNet, SOURCE_SUBJECTS, normalize_task
from event_stft.training import resolve_device, seed_everything, stable_seed


ROOT = Path(__file__).resolve().parent
TCH_CACHE = ROOT.parent / "TCH_Sequential_8ch" / "data_cache"
MODELS = ("eegnet", "shallow", "support_csp")


def task(x: np.ndarray) -> np.ndarray:
    values = x[:, :, 200:600]
    return normalize_task(values)


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


def source_arrays(source_cache: Path):
    arrays = []
    for subject in SOURCE_SUBJECTS:
        x = np.load(source_cache / "source" / f"{subject}_x.npy", allow_pickle=False)
        y = np.load(source_cache / "source" / f"{subject}_y.npy", allow_pickle=False)
        keep = y != 2
        arrays.append((task(x[keep]) if x.shape[-1] == 600 else normalize_task(x[keep]), y[keep].astype(np.int64)))
    return arrays


def make_standard(name: str):
    if name == "eegnet":
        return EEGNet(channels=8, samples=400)
    return ShallowConvNet(channels=8, samples=400, dropout=.45)


def train_standard(model, x, y, device, epochs, seed, lr=3e-4):
    model.train(); optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-3); criterion = nn.CrossEntropyLoss(label_smoothing=.05)
    generator = torch.Generator(device=device.type); generator.manual_seed(seed)
    x_t = torch.from_numpy(np.asarray(x, dtype=np.float32)).to(device); y_t = torch.from_numpy(np.asarray(y, dtype=np.int64)).to(device)
    for _ in range(epochs):
        order = torch.randperm(len(y_t), generator=generator, device=device)
        for index in order.split(min(32, len(y_t))):
            optimizer.zero_grad(set_to_none=True); loss = criterion(model(augment(x_t[index], generator)), y_t[index]); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); optimizer.step()
    return model


def evaluate_standard(model, x, y, device):
    model.eval()
    with torch.no_grad():
        prediction = model(torch.from_numpy(np.asarray(x, dtype=np.float32)).to(device)).argmax(-1).cpu().numpy()
    return float((prediction == y).mean())


def refs(x, y, support, config, seed):
    rng = np.random.default_rng(seed); index_sets = [np.asarray(support, dtype=np.int64)]
    for _ in range(config.bootstrap_references - 1):
        index_sets.append(np.concatenate([rng.choice(support[y[support] == label], size=int((y[support] == label).sum()), replace=True) for label in (0, 1)]))
    output = []
    for indices in index_sets:
        reference = CSPReference(config).fit(x[indices], y[indices])
        output.append((reference.transform(x), indices, reference))
    return output


def train_support(model, x, y, references, train_indices, device, epochs, seed):
    model.train(); optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-3); criterion = nn.CrossEntropyLoss(label_smoothing=.05)
    generator = torch.Generator(device=device.type); generator.manual_seed(seed)
    x_t = torch.from_numpy(np.asarray(x, dtype=np.float32)).to(device); y_t = torch.from_numpy(np.asarray(y, dtype=np.int64)).to(device)
    for epoch in range(epochs):
        features, _, _ = references[epoch % len(references)]
        order = torch.randperm(len(train_indices), generator=generator, device=device)
        for local in order.split(min(8, len(train_indices))):
            indices = np.asarray(train_indices)[local.cpu().numpy()]
            optimizer.zero_grad(set_to_none=True)
            logits = model(augment(x_t[indices], generator), torch.from_numpy(features[indices]).to(device))
            loss = criterion(logits, y_t[indices]); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); optimizer.step()
    return model


def evaluate_support(model, x, y, references, indices, device):
    model.eval(); values = torch.from_numpy(np.asarray(x[indices], dtype=np.float32)).to(device)
    with torch.no_grad():
        logits = torch.stack([model(values, torch.from_numpy(feature[indices]).to(device)) for feature, _, _ in references]).mean(0)
    return float((logits.argmax(-1).cpu().numpy() == y[indices]).mean())


def source_pretrain_standard(name, sources, device, epochs, seed):
    model = make_standard(name).to(device)
    for epoch in range(epochs):
        for index, (x, y) in enumerate(sources):
            train_standard(model, x, y, device, 1, stable_seed(seed, name, epoch, index))
    return model


def source_pretrain_support(sources, device, config, epochs, seed):
    model = SupportCSPShallowNet().to(device); optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-3); criterion = nn.CrossEntropyLoss(label_smoothing=.05)
    for epoch in range(epochs):
        for subject, (x, y) in enumerate(sources):
            rng = np.random.default_rng(stable_seed(seed, epoch, subject)); support=[]; query=[]
            for label in (0, 1):
                chosen = rng.permutation(np.flatnonzero(y == label)); support.extend(chosen[:12]); query.extend(chosen[12:44])
            reference = refs(x, y, np.asarray(support), config, stable_seed(seed, epoch, subject))[0][0]
            indices = np.asarray(query, dtype=np.int64); optimizer.zero_grad(set_to_none=True)
            logits = model(torch.from_numpy(x[indices]).to(device), torch.from_numpy(reference[indices]).to(device))
            loss = criterion(logits, torch.from_numpy(y[indices]).to(device)); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); optimizer.step()
    return model


def fit_and_test_fold(name, prior_state, x, y, train_indices, test_indices, device, config, seed):
    if name == "support_csp":
        model = SupportCSPShallowNet().to(device); model.load_state_dict(prior_state)
        reference = refs(x, y, train_indices, config, stable_seed(seed, "csp"))
        train_support(model, x, y, reference, train_indices, device, 25, seed)
        return evaluate_support(model, x, y, reference, test_indices, device)
    model = make_standard(name).to(device); model.load_state_dict(prior_state)
    train_standard(model, x[train_indices], y[train_indices], device, 25, seed)
    return evaluate_standard(model, x[test_indices], y[test_indices], device)


def run_subject(subject, sessions, base_states, device, config, output, independent=False):
    production = {name: copy.deepcopy(state) for name, state in base_states.items()}
    prior_references = None
    session_rows, fold_rows = [], []
    for session_number, session in enumerate(sessions, start=1):
        if independent:
            production = {name: copy.deepcopy(state) for name, state in base_states.items()}
            prior_references = None
        x = np.load(TCH_CACHE / "sessions" / f"{session}_x.npy", allow_pickle=False); y = np.load(TCH_CACHE / "sessions" / f"{session}_y.npy", allow_pickle=False); x = task(x)
        splitter = StratifiedKFold(5, shuffle=True, random_state=stable_seed(subject, session, "fold"))
        for name in MODELS:
            if name == "support_csp":
                if prior_references is None:
                    before = np.nan
                else:
                    current = SupportCSPShallowNet().to(device)
                    current.load_state_dict(production[name])
                    transferred = [(reference.transform(x), indices, reference) for _, indices, reference in prior_references]
                    before = evaluate_support(current, x, y, transferred, np.arange(len(y)), device)
            else:
                current = make_standard(name).to(device); current.load_state_dict(production[name]); before = evaluate_standard(current, x, y, device)
            scores=[]
            for fold, (train_idx, test_idx) in enumerate(splitter.split(x, y), start=1):
                score = fit_and_test_fold(name, production[name], x, y, train_idx, test_idx, device, config, stable_seed(subject, session, name, fold))
                scores.append(score); fold_rows.append({"subject_id":subject,"session_id":session,"model":name,"fold":fold,"accuracy":score,"train_trials":len(train_idx),"test_trials":len(test_idx)})
            session_rows.append({"subject_id":subject,"session_id":session,"session_order":session_number,"model":name,"before_accuracy":before,"after_cv_accuracy":float(np.mean(scores)),"after_cv_std":float(np.std(scores,ddof=1)),"trials":len(y),"left":int((y==0).sum()),"right":int((y==1).sum())})
            if not independent:
                if name == "support_csp":
                    current = SupportCSPShallowNet().to(device); current.load_state_dict(production[name]); prior_references = refs(x, y, np.arange(len(y)), config, stable_seed(subject, session, "production")); train_support(current, x, y, prior_references, np.arange(len(y)), device, 25, stable_seed(subject, session, "production")); production[name] = {key:value.detach().cpu() for key,value in current.state_dict().items()}
                else:
                    current = make_standard(name).to(device); current.load_state_dict(production[name]); train_standard(current, x, y, device, 25, stable_seed(subject, session, "production")); production[name] = {key:value.detach().cpu() for key,value in current.state_dict().items()}
        print(f"[tch] {subject} {session} complete", flush=True)
    return session_rows, fold_rows


def main():
    parser=argparse.ArgumentParser(); parser.add_argument("--source-cache", type=Path, required=True); parser.add_argument("--output", type=Path, default=ROOT / "tch_sequential_runs"); parser.add_argument("--device", default="cuda"); parser.add_argument("--source-epochs", type=int, default=8); parser.add_argument("--smoke-test", action="store_true"); parser.add_argument("--independent", action="store_true"); args=parser.parse_args()
    seed_everything(20260721); device=resolve_device(args.device); config=CSPConfig(device=args.device)
    if args.smoke_test:
        x=np.load(TCH_CACHE/"sessions"/"s1_1_x.npy",allow_pickle=False); y=np.load(TCH_CACHE/"sessions"/"s1_1_y.npy",allow_pickle=False); x=task(x); reference=refs(x,y,np.arange(20),config,1)[0][0]
        outputs={"eegnet":list(make_standard("eegnet").to(device)(torch.from_numpy(x[:2]).to(device)).shape),"shallow":list(make_standard("shallow").to(device)(torch.from_numpy(x[:2]).to(device)).shape),"support_csp":list(SupportCSPShallowNet().to(device)(torch.from_numpy(x[:2]).to(device),torch.from_numpy(reference[:2]).to(device)).shape)}
        print({"shape":list(x.shape),"labels":np.bincount(y).tolist(),"outputs":outputs},flush=True); return
    sources=source_arrays(args.source_cache); args.output.mkdir(parents=True, exist_ok=True)
    checkpoints=args.output/"source_checkpoints.pt"
    if checkpoints.exists(): base_states=torch.load(checkpoints,map_location="cpu",weights_only=False)
    else:
        base_states={"eegnet":source_pretrain_standard("eegnet",sources,device,args.source_epochs,1).cpu().state_dict(),"shallow":source_pretrain_standard("shallow",sources,device,args.source_epochs,2).cpu().state_dict(),"support_csp":source_pretrain_support(sources,device,config,30,3).cpu().state_dict()}; torch.save(base_states,checkpoints)
    all_sessions=[]; all_folds=[]
    for subject,sessions in (("s1",["s1_1","s1_2","s1_3","s1_4","s1_5"]),("s2",["s2_1","s2_2","s2_3","s2_4"])):
        rows,folds=run_subject(subject,sessions,base_states,device,config,args.output,args.independent); all_sessions.extend(rows); all_folds.extend(folds); pd.DataFrame(all_sessions).to_csv(args.output/"session_results.csv",index=False,encoding="utf-8-sig"); pd.DataFrame(all_folds).to_csv(args.output/"fold_results.csv",index=False,encoding="utf-8-sig")
    protocol = "each session independently starts from the same source base; after=5-fold current-session adaptation" if args.independent else "before=current production model on all new-session trials; after=5-fold current-session adaptation; production then updates on all current-session trials"
    (args.output/"run_config.json").write_text(json.dumps({"models":MODELS,"source_epochs_standard":args.source_epochs,"source_epochs_support":30,"target_epochs_per_update":25,"protocol":protocol},ensure_ascii=False,indent=2),encoding="utf-8")


if __name__=="__main__": main()
