"""Support-conditioned shrinkage CSP + ShallowConvNet for few-shot MI."""
from __future__ import annotations

import json
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Dict, Iterable, List, Mapping, Sequence, Tuple

import numpy as np
import pandas as pd
import torch
import torch.nn as nn
from scipy.linalg import eigh
from scipy.signal import butter, sosfiltfilt
from scipy.stats import wilcoxon
from sklearn.model_selection import StratifiedKFold

from .data import LABELED_SUBJECTS, SUBSAMPLE_SEEDS, load_fold_assignments, load_subject
from .models import ShallowConvNet, count_parameters
from .training import TrainConfig, _scheduler_lambda, augment_waveforms, make_loader, prepare_model_input, resolve_device, seed_everything, stable_seed, train_fold

MODEL_SHALLOW = "global_shallow"
MODEL_CSP_SCRATCH = "support_csp_shallow_scratch"
MODEL_CSP_EPISODIC = "support_csp_shallow_episodic"
MODELS = (MODEL_SHALLOW, MODEL_CSP_SCRATCH, MODEL_CSP_EPISODIC)
SOURCE_SUBJECTS = tuple(f"sub{i}" for i in range(16, 50))


@dataclass(frozen=True)
class CSPConfig:
    device: str = "cuda"
    seed: int = 20260721
    all_trial_seeds: Tuple[int, ...] = (20260720, 20260721, 20260722)
    bands: Tuple[Tuple[float, float], ...] = ((8., 13.), (13., 20.), (20., 30.))
    filters_per_side: int = 2
    shrinkage: float = .15
    source_epochs: int = 30
    source_support_per_class: int = 12
    source_query_per_class: int = 32
    source_lr: float = 3e-4
    target_epochs: int = 150
    target_patience: int = 20
    batch_size: int = 32
    fewshot_batch_size: int = 8
    finetune_lr: float = 3e-4
    weight_decay: float = 1e-3
    label_smoothing: float = .05
    warmup_epochs: int = 5
    amplitude_low: float = .95
    amplitude_high: float = 1.05
    max_shift_samples: int = 10
    bootstrap_references: int = 8
    freeze_wave_epochs: int = 10
    l2sp_lambda: float = 1e-4
    crossfit_train_queries: bool = False

    def shallow(self) -> TrainConfig:
        return TrainConfig(device=self.device, seed=self.seed, all_trial_seeds=self.all_trial_seeds,
            max_epochs=self.target_epochs, patience=self.target_patience, batch_size_all=self.batch_size,
            batch_size_fewshot=self.fewshot_batch_size, learning_rate=self.finetune_lr,
            weight_decay=self.weight_decay, label_smoothing=self.label_smoothing,
            warmup_epochs=self.warmup_epochs, amplitude_low=self.amplitude_low,
            amplitude_high=self.amplitude_high, max_shift_samples=self.max_shift_samples)


def normalize_task(x: np.ndarray) -> np.ndarray:
    x = np.asarray(x, dtype=np.float32)
    return ((x - x.mean(axis=-1, keepdims=True)) / np.maximum(x.std(axis=-1, keepdims=True), 1e-6)).astype(np.float32)


class CSPReference:
    """Train-support-only discriminative spatial filters and class prototypes."""
    def __init__(self, config: CSPConfig) -> None:
        self.config=config; self.filters=[]; self.feature_mean=None; self.feature_sd=None; self.prototypes=None
        self.sos=[butter(4, band, btype="bandpass", fs=200, output="sos") for band in config.bands]

    def _band(self, x: np.ndarray, sos: np.ndarray) -> np.ndarray:
        return sosfiltfilt(sos, x.astype(np.float64), axis=-1).astype(np.float32)

    def _covariance(self, values: np.ndarray) -> np.ndarray:
        cov=np.einsum("nct,ndt->ncd",values,values)/values.shape[-1]
        trace=np.trace(cov,axis1=1,axis2=2)[:,None,None]
        return cov/np.maximum(trace,1e-8)

    def fit(self, x: np.ndarray, y: np.ndarray) -> "CSPReference":
        if min((y==0).sum(),(y==1).sum()) < 4: raise ValueError("CSP support needs at least four trials/class")
        features=[]
        for sos in self.sos:
            band=self._band(x,sos); cov=self._covariance(band)
            covs=[]
            for label in (0,1):
                value=cov[y==label].mean(axis=0); identity=np.eye(value.shape[0])*np.trace(value)/value.shape[0]
                covs.append((1-self.config.shrinkage)*value+self.config.shrinkage*identity)
            values,vectors=eigh(covs[0],covs[0]+covs[1]); indices=np.r_[np.arange(self.config.filters_per_side),np.arange(-self.config.filters_per_side,0)]
            matrix=vectors[:,indices].T.astype(np.float32); self.filters.append(matrix)
            projected=np.einsum("fc,nct->nft",matrix,band); features.append(np.log(np.maximum(projected.var(axis=-1),1e-8)))
        raw=np.concatenate(features,axis=1).astype(np.float32)
        self.feature_mean=raw.mean(axis=0); self.feature_sd=np.maximum(raw.std(axis=0),1e-6); standardized=(raw-self.feature_mean)/self.feature_sd
        self.prototypes=np.stack([standardized[y==label].mean(axis=0) for label in (0,1)]).astype(np.float32)
        return self

    def transform(self, x: np.ndarray) -> np.ndarray:
        if self.feature_mean is None: raise RuntimeError("CSPReference must be fitted first")
        features=[]
        for sos,matrix in zip(self.sos,self.filters):
            projected=np.einsum("fc,nct->nft",matrix,self._band(x,sos)); features.append(np.log(np.maximum(projected.var(axis=-1),1e-8)))
        standardized=(np.concatenate(features,axis=1)-self.feature_mean)/self.feature_sd
        distances=np.linalg.norm(standardized[:,None,:]-self.prototypes[None,:,:],axis=-1)
        return np.concatenate([standardized,distances],axis=1).astype(np.float32)


class SupportCSPShallowNet(nn.Module):
    def __init__(self, csp_dim: int=14, dropout: float=.45) -> None:
        super().__init__(); base=ShallowConvNet(channels=8,samples=400,dropout=dropout)
        self.wave_features=base.features; self.wave_dim=840
        self.wave_head=nn.Sequential(nn.LayerNorm(self.wave_dim),nn.Linear(self.wave_dim,96),nn.ELU(),nn.Dropout(dropout))
        self.csp_head=nn.Sequential(nn.LayerNorm(csp_dim),nn.Linear(csp_dim,32),nn.ELU())
        self.classifier=nn.Sequential(nn.Linear(128,64),nn.ELU(),nn.Dropout(dropout),nn.Linear(64,2))
    def forward(self, x: torch.Tensor, csp: torch.Tensor, return_features: bool=False):
        if x.ndim==3: x=x.unsqueeze(1)
        wave=self.wave_features(x).flatten(1); features=torch.cat([self.wave_head(wave),self.csp_head(csp)],dim=1); logits=self.classifier(features)
        return (logits,features) if return_features else logits


def _model_inputs(raw: torch.Tensor, train: bool, config: CSPConfig) -> torch.Tensor:
    values=raw[...,200:600]; values=(values-values.mean(-1,keepdim=True))/values.std(-1,keepdim=True,unbiased=False).clamp_min(1e-6)
    return augment_waveforms(values,config.shallow()) if train else values


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


def _evaluate(model, loader, csp_sets, start_indices, config, device):
    model.eval(); labels=[]; pred=[]; losses=[]; criterion=nn.CrossEntropyLoss()
    with torch.no_grad():
        for raw,target,index in loader:
            raw,target,index=raw.to(device),target.to(device),index.numpy(); inputs=_model_inputs(raw,False,config)
            logits=torch.stack([model(inputs,torch.from_numpy(values[start_indices[index]]).to(device)) for values in csp_sets]).mean(0)
            losses.append(float(criterion(logits,target))*len(target)); labels.append(target.cpu().numpy()); pred.append(logits.argmax(-1).cpu().numpy())
    y=np.concatenate(labels); p=np.concatenate(pred)
    return {"loss":sum(losses)/len(y),"accuracy":float((p==y).mean()),"balanced_accuracy":float(np.mean([(p[y==c]==c).mean() for c in (0,1)])),"macro_f1":float(__import__('sklearn.metrics').metrics.f1_score(y,p,average='macro'))}


def _indexed_loader(x,y,indices,batch,shuffle,seed):
    dataset=torch.utils.data.TensorDataset(torch.from_numpy(np.asarray(x[indices],dtype=np.float32)),torch.from_numpy(np.asarray(y[indices],dtype=np.int64)),torch.from_numpy(np.arange(len(indices),dtype=np.int64)))
    generator=torch.Generator(); generator.manual_seed(seed); return torch.utils.data.DataLoader(dataset,batch_size=min(batch,len(dataset)),shuffle=shuffle,generator=generator,num_workers=0)


def train_target_fold(x,y,split,mode,seed,config,initial_state=None,evaluate_test=True,checkpoint=None):
    seed_everything(seed); device=resolve_device(config.device); task=normalize_task(x[:, :, 200:600]); csp_sets=_bootstrap_csp(task,y,split['train'],config,stable_seed(seed,'bootstrap'))
    model=SupportCSPShallowNet(csp_sets[0].shape[1]).to(device)
    if initial_state is not None: model.load_state_dict(initial_state,strict=True)
    batch=config.fewshot_batch_size if sum(len(v) for v in split.values())<=40 else config.batch_size
    loaders={role:_indexed_loader(x,y,indices,batch,role=='train',stable_seed(seed,role)) for role,indices in split.items() if role!='test' or evaluate_test}
    train_csp=None
    if config.crossfit_train_queries:
        train_csp=np.zeros_like(csp_sets[0])
        local_y=y[split['train']]
        for _, query_local in StratifiedKFold(3,shuffle=True,random_state=stable_seed(seed,'crossfit')).split(split['train'],local_y):
            query_indices=split['train'][query_local]; support=np.setdiff1d(split['train'],query_indices,assume_unique=True)
            query_sets=_bootstrap_csp(task,y,support,config,stable_seed(seed,'crossfit',len(query_indices)))
            train_csp[query_indices]=np.mean([values[query_indices] for values in query_sets],axis=0)
    optimizer=torch.optim.AdamW(model.parameters(),lr=config.finetune_lr,weight_decay=config.weight_decay); scheduler=torch.optim.lr_scheduler.LambdaLR(optimizer,lambda e:_scheduler_lambda(e,config.target_epochs,config.warmup_epochs)); criterion=nn.CrossEntropyLoss(label_smoothing=config.label_smoothing); amp=device.type=='cuda'; scaler=torch.amp.GradScaler(device.type,enabled=amp)
    l2sp={name:value.to(device) for name,value in initial_state.items() if name.startswith('wave_features')} if initial_state is not None else {}
    best=None; best_score=-1.; best_loss=float('inf'); best_epoch=0; stale=0
    for epoch in range(1,config.target_epochs+1):
        freeze=initial_state is not None and epoch<=config.freeze_wave_epochs
        for parameter in model.wave_features.parameters(): parameter.requires_grad=not freeze
        model.train()
        for raw,target,index in loaders['train']:
            raw,target=raw.to(device),target.to(device); indices=split['train'][index.numpy()]; csp_batch=torch.from_numpy((train_csp if train_csp is not None else csp_sets[(epoch-1)%len(csp_sets)])[indices]).to(device); optimizer.zero_grad(set_to_none=True)
            with torch.autocast(device_type=device.type,enabled=amp): loss=criterion(model(_model_inputs(raw,True,config),csp_batch),target)
            if l2sp and not freeze:
                penalty=sum((parameter-l2sp[name]).square().sum() for name,parameter in model.named_parameters() if name in l2sp)
                loss=loss+config.l2sp_lambda*penalty
            scaler.scale(loss).backward(); scaler.unscale_(optimizer); torch.nn.utils.clip_grad_norm_(model.parameters(),1.0); scaler.step(optimizer); scaler.update()
        val=_evaluate(model,loaders['val'],csp_sets,split['val'],config,device); scheduler.step(); good=val['balanced_accuracy']>best_score+1e-8 or (abs(val['balanced_accuracy']-best_score)<1e-8 and val['loss']<best_loss)
        if good: best={k:v.detach().cpu().clone() for k,v in model.state_dict().items()}; best_score=val['balanced_accuracy']; best_loss=val['loss']; best_epoch=epoch; stale=0
        else: stale+=1
        if stale>=config.target_patience: break
    model.load_state_dict(best); result={"model":mode,"best_epoch":best_epoch,"parameter_count":count_parameters(model),"val_balanced_accuracy":best_score}
    if evaluate_test: result.update(_evaluate(model,loaders['test'],csp_sets,split['test'],config,device))
    if checkpoint: checkpoint.parent.mkdir(parents=True,exist_ok=True); torch.save({"state":best,"csp_dim":csp_sets[0].shape[1],"result":result,"config":asdict(config)},checkpoint)
    return result


def source_pretrain(source_cache: Path, run_dir: Path, config: CSPConfig) -> Path:
    checkpoint=run_dir/'source_pretrained.pt'; history_path=run_dir/'source_history.csv'
    if checkpoint.exists(): return checkpoint
    device=resolve_device(config.device); seed_everything(config.seed); model=SupportCSPShallowNet().to(device); opt=torch.optim.AdamW(model.parameters(),lr=config.source_lr,weight_decay=config.weight_decay); criterion=nn.CrossEntropyLoss(label_smoothing=config.label_smoothing); amp=device.type=='cuda'; scaler=torch.amp.GradScaler(device.type,enabled=amp); rows=[]; rng=np.random.default_rng(config.seed)
    subjects=[]
    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; subjects.append((normalize_task(x[keep]),y[keep]))
    for epoch in range(1,config.source_epochs+1):
        model.train(); losses=[]
        for x,y in subjects:
            support=[]; query=[]
            for label in (0,1):
                pool=np.flatnonzero(y==label); chosen=rng.permutation(pool); support.extend(chosen[:config.source_support_per_class]); query.extend(chosen[config.source_support_per_class:config.source_support_per_class+config.source_query_per_class])
            reference=CSPReference(config).fit(x[np.asarray(support)],y[np.asarray(support)]); csp=reference.transform(x[np.asarray(query)]); raw=torch.from_numpy(x[np.asarray(query)]).to(device); target=torch.from_numpy(y[np.asarray(query)]).to(device); csp_t=torch.from_numpy(csp).to(device); opt.zero_grad(set_to_none=True)
            with torch.autocast(device_type=device.type,enabled=amp): loss=criterion(model(raw,csp_t),target)
            scaler.scale(loss).backward(); scaler.unscale_(opt); torch.nn.utils.clip_grad_norm_(model.parameters(),1.0); scaler.step(opt); scaler.update(); losses.append(float(loss.detach()))
        rows.append({"epoch":epoch,"loss":float(np.mean(losses))}); print(f'[csp source] epoch={epoch} loss={rows[-1]["loss"]:.4f}',flush=True)
    run_dir.mkdir(parents=True,exist_ok=True); pd.DataFrame(rows).to_csv(history_path,index=False,encoding='utf-8-sig'); torch.save({"state":model.state_dict(),"config":asdict(config)},checkpoint); return checkpoint


def protocols(config):
    for seed in config.all_trial_seeds: yield 'all','all',seed
    for seed in SUBSAMPLE_SEEDS: yield '20x20',f'20_seed_{seed}',stable_seed('20csp',seed)


def run_evaluation(cache, source_checkpoint, run_dir, config, max_new=None):
    out=run_dir/'final'; path=out/'fold_results.csv'; rows=pd.read_csv(path).to_dict('records') if path.exists() else []; done={(r['model'],r['protocol'],r['split_set'],int(r['train_seed']),r['subject_id'],int(r['fold'])) for r in rows}; assignments=load_fold_assignments(cache/'fold_assignments.csv'); state=torch.load(source_checkpoint,map_location='cpu',weights_only=False)['state']; made=0
    for protocol,split_set,run_seed in protocols(config):
        seeds=config.all_trial_seeds if protocol=='all' else (run_seed,)
        for subject in LABELED_SUBJECTS:
            x,y=load_subject(cache,subject)
            for fold in range(1,6):
                split=assignments[(split_set,subject,fold)]
                for mode,initial in ((MODEL_CSP_SCRATCH,None),(MODEL_CSP_EPISODIC,state)):
                    for train_seed in seeds:
                        key=(mode,protocol,split_set,train_seed,subject,fold)
                        if key in done: continue
                        if max_new is not None and made>=max_new: pd.DataFrame(rows).to_csv(path,index=False,encoding='utf-8-sig'); return pd.DataFrame(rows)
                        print(f'[csp final] {mode} {protocol} {subject} {split_set} fold={fold}',flush=True)
                        result=train_target_fold(x,y,split,mode,stable_seed(mode,protocol,split_set,train_seed,subject,fold),config,initial,True,out/'checkpoints'/f'{mode}_{protocol}_{split_set}_{train_seed}_{subject}_{fold}.pt')
                        rows.append({"protocol":protocol,"split_set":split_set,"train_seed":train_seed,"subject_id":subject,"fold":fold,**result}); done.add(key); made+=1; out.mkdir(parents=True,exist_ok=True); pd.DataFrame(rows).to_csv(path,index=False,encoding='utf-8-sig')
    frame=pd.DataFrame(rows); subject=frame.groupby(['protocol','model','subject_id'],as_index=False).agg(mean_accuracy=('accuracy','mean'),std_accuracy=('accuracy','std'),mean_balanced_accuracy=('balanced_accuracy','mean')); group=subject.groupby(['protocol','model'],as_index=False).agg(mean_accuracy=('mean_accuracy','mean'),std_subject_accuracy=('mean_accuracy','std'),mean_balanced_accuracy=('mean_balanced_accuracy','mean'),subjects=('subject_id','count')); subject.to_csv(out/'subject_summary.csv',index=False,encoding='utf-8-sig'); group.to_csv(out/'group_summary.csv',index=False,encoding='utf-8-sig'); (out/'run_config.json').write_text(json.dumps(asdict(config),ensure_ascii=False,indent=2),encoding='utf-8'); return frame


def smoke_test():
    config=CSPConfig(device='cpu'); x=normalize_task(np.random.default_rng(1).normal(size=(12,8,400)).astype(np.float32)); y=np.array([0]*6+[1]*6); ref=CSPReference(config).fit(x,y); features=ref.transform(x); model=SupportCSPShallowNet(features.shape[1]); logits=model(torch.from_numpy(x[:2]),torch.from_numpy(features[:2])); return {"csp_shape":list(features.shape),"logits":list(logits.shape),"parameters":count_parameters(model)}
