"""Support-conditioned CSP Proto-FiLM ShallowNet for few-shot MI adaptation."""
from __future__ import annotations

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

import numpy as np
import pandas as pd
import torch
import torch.nn as nn

from .data import LABELED_SUBJECTS, SUBSAMPLE_SEEDS, load_fold_assignments, load_subject
from .models import ShallowConvNet, count_parameters
from .support_csp import CSPConfig, CSPReference, SOURCE_SUBJECTS, normalize_task
from .training import _scheduler_lambda, augment_waveforms, resolve_device, seed_everything, stable_seed


MODEL_PROTOFILM_SCRATCH = "support_csp_protofilm_scratch"
MODEL_PROTOFILM_EPISODIC = "support_csp_protofilm_episodic"


class SupportCSPProtoFiLMNet(nn.Module):
    """Fuse CSP features with a support-conditioned ShallowConvNet embedding."""

    def __init__(self, csp_dim: int = 14, context_dim: int = 48, dropout: float = 0.45) -> None:
        super().__init__()
        base = ShallowConvNet(channels=8, samples=400, dropout=dropout)
        self.wave_features = base.features
        self.wave_dim = 840
        self.wave_norm = nn.LayerNorm(self.wave_dim)
        self.wave_linear = nn.Linear(self.wave_dim, 96)
        self.wave_dropout = nn.Dropout(dropout)
        self.context_head = nn.Sequential(
            nn.LayerNorm(context_dim),
            nn.Linear(context_dim, 128),
            nn.ELU(),
            nn.Linear(128, 192),
        )
        self.csp_head = nn.Sequential(nn.LayerNorm(csp_dim), nn.Linear(csp_dim, 32), nn.ELU())
        self.classifier = nn.Sequential(
            nn.Linear(130, 64), nn.ELU(), nn.Dropout(dropout), nn.Linear(64, 2)
        )

    def _film(self, context: torch.Tensor, batch: int) -> Tuple[torch.Tensor, torch.Tensor]:
        if context.ndim == 1:
            context = context.unsqueeze(0)
        if context.shape[0] == 1:
            context = context.expand(batch, -1)
        gamma, beta = self.context_head(context).chunk(2, dim=-1)
        return 1.0 + 0.2 * torch.tanh(gamma), 0.2 * torch.tanh(beta)

    def _wave_embedding(
        self, x: torch.Tensor, gamma: torch.Tensor, beta: torch.Tensor, dropout: bool
    ) -> torch.Tensor:
        if x.ndim == 3:
            x = x.unsqueeze(1)
        values = self.wave_features(x).flatten(1)
        values = torch.nn.functional.elu(self.wave_linear(self.wave_norm(values)))
        values = gamma * values + beta
        return self.wave_dropout(values) if dropout else values

    def forward(
        self,
        query_x: torch.Tensor,
        query_csp: torch.Tensor,
        support_x: torch.Tensor,
        support_y: torch.Tensor,
        context: torch.Tensor,
        return_features: bool = False,
    ):
        batch = query_x.shape[0]
        gamma, beta = self._film(context, batch)
        query_wave = self._wave_embedding(query_x, gamma, beta, dropout=self.training)
        support_gamma = gamma[:1].expand(support_x.shape[0], -1)
        support_beta = beta[:1].expand(support_x.shape[0], -1)
        support_wave = self._wave_embedding(support_x, support_gamma, support_beta, dropout=False)
        prototypes = torch.stack([support_wave[support_y == label].mean(0) for label in (0, 1)])
        raw_distances = torch.linalg.vector_norm(
            query_wave[:, None, :] - prototypes[None, :, :], dim=-1
        ) / np.sqrt(query_wave.shape[-1])
        features = torch.cat([query_wave, self.csp_head(query_csp), raw_distances], 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 _context(reference: CSPReference) -> np.ndarray:
    left, right = reference.prototypes
    return np.concatenate([left, right, left - right, np.abs(left - right)]).astype(np.float32)


def _references(
    task: np.ndarray, y: np.ndarray, support: np.ndarray, config: CSPConfig, seed: int
) -> List[Dict[str, np.ndarray]]:
    rng = np.random.default_rng(seed)
    picked_sets = [np.asarray(support, dtype=np.int64)]
    for _ in range(config.bootstrap_references - 1):
        picked_sets.append(
            np.concatenate(
                [
                    rng.choice(support[y[support] == label], size=int((y[support] == label).sum()), replace=True)
                    for label in (0, 1)
                ]
            ).astype(np.int64)
        )
    references = []
    for picked in picked_sets:
        reference = CSPReference(config).fit(task[picked], y[picked])
        references.append(
            {
                "features": reference.transform(task),
                "context": _context(reference),
                "support_x": task[picked],
                "support_y": y[picked],
            }
        )
    return references


def _loader(x: np.ndarray, y: np.ndarray, indices: np.ndarray, batch: int, shuffle: bool, seed: int):
    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().manual_seed(seed)
    return torch.utils.data.DataLoader(dataset, batch_size=min(batch, len(dataset)), shuffle=shuffle, generator=generator, num_workers=0)


def _ref_tensors(reference: Mapping[str, np.ndarray], device: torch.device):
    return (
        torch.from_numpy(reference["support_x"]).to(device),
        torch.from_numpy(reference["support_y"]).to(device),
        torch.from_numpy(reference["context"]).to(device),
    )


def _evaluate(model, loader, references, indices: np.ndarray, config: CSPConfig, device: torch.device):
    model.eval()
    labels: List[np.ndarray] = []
    predictions: List[np.ndarray] = []
    losses: List[float] = []
    criterion = nn.CrossEntropyLoss()
    tensors = [_ref_tensors(reference, device) for reference in references]
    with torch.no_grad():
        for raw, target, local_index in loader:
            raw, target = raw.to(device), target.to(device)
            global_index = indices[local_index.numpy()]
            query = _model_inputs(raw, False, config)
            logits = torch.stack(
                [
                    model(
                        query,
                        torch.from_numpy(reference["features"][global_index]).to(device),
                        support_x,
                        support_y,
                        context,
                    )
                    for reference, (support_x, support_y, context) in zip(references, tensors)
                ]
            ).mean(0)
            losses.append(float(criterion(logits, target)) * len(target))
            labels.append(target.cpu().numpy())
            predictions.append(logits.argmax(-1).cpu().numpy())
    y = np.concatenate(labels)
    p = np.concatenate(predictions)
    return {
        "loss": sum(losses) / len(y),
        "accuracy": float((p == y).mean()),
        "balanced_accuracy": float(np.mean([(p[y == label] == label).mean() for label in (0, 1)])),
        "macro_f1": float(__import__("sklearn.metrics").metrics.f1_score(y, p, average="macro")),
    }


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])
    references = _references(task, y, split["train"], config, stable_seed(seed, "bootstrap"))
    model = SupportCSPProtoFiLMNet(csp_dim=references[0]["features"].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(values) for values in split.values()) <= 40 else config.batch_size
    loaders = {
        role: _loader(x, y, indices, batch, role == "train", stable_seed(seed, role))
        for role, indices in split.items()
        if role != "test" or evaluate_test
    }
    optimizer = torch.optim.AdamW(model.parameters(), lr=config.finetune_lr, weight_decay=config.weight_decay)
    scheduler = torch.optim.lr_scheduler.LambdaLR(
        optimizer, lambda epoch: _scheduler_lambda(epoch, 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 else {}
    best = None
    best_score, best_loss, best_epoch, stale = -1.0, float("inf"), 0, 0
    for epoch in range(1, config.target_epochs + 1):
        frozen = initial_state is not None and epoch <= config.freeze_wave_epochs
        for parameter in model.wave_features.parameters():
            parameter.requires_grad = not frozen
        model.train()
        reference = references[(epoch - 1) % len(references)]
        support_x, support_y, context = _ref_tensors(reference, device)
        for raw, target, local_index in loaders["train"]:
            raw, target = raw.to(device), target.to(device)
            global_index = split["train"][local_index.numpy()]
            csp = torch.from_numpy(reference["features"][global_index]).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, support_x, support_y, context), target)
            if l2sp and not frozen:
                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"], references, 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 = {name: value.detach().cpu().clone() for name, value in model.state_dict().items()}
            best_score, best_loss, best_epoch, stale = val["balanced_accuracy"], val["loss"], epoch, 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"], references, split["test"], config, device))
    if checkpoint:
        checkpoint.parent.mkdir(parents=True, exist_ok=True)
        torch.save({"state": best, "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"
    if checkpoint.exists():
        return checkpoint
    seed_everything(config.seed)
    device = resolve_device(config.device)
    model = SupportCSPProtoFiLMNet().to(device)
    optimizer = 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)
    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
        # Source caches retain the -1..2 s epoch; source and target must both use 0..2 s.
        task = x[keep, :, 200:600] if x.shape[-1] == 600 else x[keep]
        subjects.append((subject, normalize_task(task), y[keep]))
    history = []
    for epoch in range(1, config.source_epochs + 1):
        model.train()
        losses = []
        for subject, task, y in subjects:
            support, query = [], []
            for label in (0, 1):
                chosen = rng.permutation(np.flatnonzero(y == label))
                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])
            support = np.asarray(support, dtype=np.int64)
            query = np.asarray(query, dtype=np.int64)
            reference = _references(task, y, support, config, stable_seed(config.seed, epoch, subject))[0]
            support_x, support_y, context = _ref_tensors(reference, device)
            raw = torch.from_numpy(task[query]).to(device)
            target = torch.from_numpy(y[query]).to(device)
            csp = torch.from_numpy(reference["features"][query]).to(device)
            optimizer.zero_grad(set_to_none=True)
            with torch.autocast(device_type=device.type, enabled=amp):
                loss = criterion(model(raw, csp, support_x, support_y, context), target)
            scaler.scale(loss).backward()
            scaler.unscale_(optimizer)
            torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
            scaler.step(optimizer)
            scaler.update()
            losses.append(float(loss.detach()))
        history.append({"epoch": epoch, "loss": float(np.mean(losses))})
        print(f"[protofilm source] epoch={epoch} loss={history[-1]['loss']:.4f}", flush=True)
    run_dir.mkdir(parents=True, exist_ok=True)
    pd.DataFrame(history).to_csv(run_dir / "source_history.csv", index=False, encoding="utf-8-sig")
    torch.save({"state": model.state_dict(), "config": asdict(config)}, checkpoint)
    return checkpoint


def _protocols(config: CSPConfig, selected: Sequence[str]):
    if "all" in selected:
        for seed in config.all_trial_seeds:
            yield "all", "all", seed
    if "20x20" in selected:
        for seed in SUBSAMPLE_SEEDS:
            yield "20x20", f"20_seed_{seed}", stable_seed("20protofilm", seed)


def run_evaluation(cache: Path, source_checkpoint: Path, run_dir: Path, config: CSPConfig, protocols=("all", "20x20"), max_new=None):
    output = run_dir / "final"
    path = output / "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, protocols):
        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_PROTOFILM_SCRATCH, None), (MODEL_PROTOFILM_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"[protofilm] {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, output / "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
                        output.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")
    )
    output.mkdir(parents=True, exist_ok=True)
    subject.to_csv(output / "subject_summary.csv", index=False, encoding="utf-8-sig")
    group.to_csv(output / "group_summary.csv", index=False, encoding="utf-8-sig")
    (output / "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")
    rng = np.random.default_rng(1)
    task = normalize_task(rng.normal(size=(24, 8, 400)).astype(np.float32))
    y = np.array([0] * 12 + [1] * 12)
    reference = _references(task, y, np.arange(24), config, 1)[0]
    model = SupportCSPProtoFiLMNet(csp_dim=reference["features"].shape[1])
    logits = model(
        torch.from_numpy(task[:2]), torch.from_numpy(reference["features"][:2]),
        torch.from_numpy(reference["support_x"]), torch.from_numpy(reference["support_y"]), torch.from_numpy(reference["context"]),
    )
    return {"logits": list(logits.shape), "parameters": count_parameters(model), "context": list(reference["context"].shape)}
