"""Train-only KMeans hard routing with standard ShallowConvNet experts."""

from __future__ import annotations

import copy
import json
import math
from dataclasses import asdict, dataclass
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 scipy.signal import welch
from scipy.stats import wilcoxon
from sklearn.cluster import KMeans
from sklearn.decomposition import PCA
from sklearn.metrics import (
    accuracy_score,
    balanced_accuracy_score,
    confusion_matrix,
    f1_score,
    roc_auc_score,
    silhouette_score,
)
from sklearn.preprocessing import StandardScaler

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,
    evaluate_loader,
    make_loader,
    prepare_model_input,
    resolve_device,
    seed_everything,
    stable_seed,
    train_fold,
)


MODEL_GLOBAL = "global_shallow"
MODEL_SCRATCH = "hard_k2_scratch"
MODEL_GLOBAL_INIT = "hard_k2_global_init"
MODEL_NAMES = (MODEL_GLOBAL, MODEL_SCRATCH, MODEL_GLOBAL_INIT)


@dataclass(frozen=True)
class HardRoutingConfig:
    device: str = "cuda"
    seed: int = 20260720
    all_trial_seeds: Tuple[int, ...] = (20260720, 20260721, 20260722)
    kmeans_clusters: int = 2
    kmeans_n_init: int = 50
    pca_components: int = 8
    scratch_epochs: int = 150
    scratch_patience: int = 20
    global_init_epochs: int = 50
    global_init_patience: int = 10
    scratch_lr: float = 3e-4
    global_init_lr: float = 1e-4
    weight_decay: float = 1e-3
    label_smoothing: float = 0.05
    warmup_epochs: int = 5
    batch_size_all: int = 32
    batch_size_fewshot: int = 8
    gradient_clip: float = 1.0
    amplitude_low: float = 0.95
    amplitude_high: float = 1.05
    max_shift_samples: int = 10

    def training_config(self) -> TrainConfig:
        return TrainConfig(
            device=self.device,
            seed=self.seed,
            all_trial_seeds=self.all_trial_seeds,
            max_epochs=self.scratch_epochs,
            patience=self.scratch_patience,
            batch_size_all=self.batch_size_all,
            batch_size_fewshot=self.batch_size_fewshot,
            learning_rate=self.scratch_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,
            gradient_clip=self.gradient_clip,
        )


def baseline_features(x: np.ndarray, sample_rate: int = 200, baseline_samples: int = 200) -> np.ndarray:
    """Label-free trial-state features computed only from event-preceding baseline."""
    baseline = np.asarray(x[:, :, :baseline_samples], dtype=np.float64)
    baseline = baseline - baseline.mean(axis=-1, keepdims=True)
    frequencies, power = welch(baseline, fs=sample_rate, nperseg=128, axis=-1)
    epsilon = 1e-12
    total_mask = (frequencies >= 1.0) & (frequencies <= 40.0)
    total = power[:, :, total_mask].mean(axis=-1) + epsilon
    relative_bands = []
    for low, high in ((1, 4), (4, 8), (8, 13), (13, 20), (20, 30), (30, 40)):
        mask = (frequencies >= low) & (frequencies <= high)
        band = power[:, :, mask].mean(axis=-1) + epsilon
        relative_bands.append(np.log(band / total))
    correlations = np.asarray([np.corrcoef(trial) for trial in baseline])
    correlations = np.nan_to_num(correlations, nan=0.0, posinf=0.0, neginf=0.0)
    triangle = np.triu_indices(baseline.shape[1], k=1)
    values = np.concatenate(
        [
            np.log(baseline.std(axis=-1) + epsilon),
            *relative_bands,
            correlations[:, triangle[0], triangle[1]],
        ],
        axis=1,
    )
    if not np.isfinite(values).all():
        raise RuntimeError("Non-finite baseline routing features")
    return values.astype(np.float32)


class TrainOnlyKMeansRouter:
    def __init__(self, config: HardRoutingConfig) -> None:
        self.config = config
        self.scaler = StandardScaler()
        self.pca: PCA | None = None
        self.kmeans = KMeans(
            n_clusters=config.kmeans_clusters,
            n_init=config.kmeans_n_init,
            random_state=config.seed,
        )
        self.fitted_samples = 0

    def fit(self, features: np.ndarray) -> "TrainOnlyKMeansRouter":
        if len(features) < self.config.kmeans_clusters * 2:
            raise ValueError("Insufficient train samples for KMeans routing")
        standardized = self.scaler.fit_transform(features)
        components = min(self.config.pca_components, standardized.shape[1], len(features) - 2)
        if components < 1:
            raise ValueError("PCA requires at least one component")
        self.pca = PCA(n_components=components, whiten=True, random_state=self.config.seed)
        embedded = self.pca.fit_transform(standardized)
        self.kmeans.fit(embedded)
        self.fitted_samples = len(features)
        return self

    @property
    def pca_components(self) -> int:
        if self.pca is None:
            raise RuntimeError("Router has not been fitted")
        return int(self.pca.n_components_)

    def transform(self, features: np.ndarray) -> np.ndarray:
        if self.pca is None:
            raise RuntimeError("Router has not been fitted")
        return self.pca.transform(self.scaler.transform(features))

    def predict(self, features: np.ndarray) -> np.ndarray:
        return self.kmeans.predict(self.transform(features)).astype(np.int64)

    def train_silhouette(self, features: np.ndarray) -> float:
        embedded = self.transform(features)
        labels = self.kmeans.predict(embedded)
        counts = np.bincount(labels, minlength=self.config.kmeans_clusters)
        if (counts < 2).any():
            return float("nan")
        return float(silhouette_score(embedded, labels))


def class_counts(labels: np.ndarray, clusters: np.ndarray, cluster: int) -> Tuple[int, int]:
    selected = labels[clusters == cluster]
    return int((selected == 0).sum()), int((selected == 1).sum())


def fallback_reason(train_labels: np.ndarray, validation_labels: np.ndarray) -> str:
    if len(train_labels) == 0:
        return "empty_train_cluster"
    if len(np.unique(train_labels)) < 2:
        return "missing_train_class"
    if len(validation_labels) == 0:
        return "empty_validation_cluster"
    return ""


def _criterion(labels: np.ndarray, config: HardRoutingConfig, device: torch.device) -> nn.Module:
    counts = np.bincount(labels, minlength=2).astype(np.float64)
    if (counts == 0).any():
        raise ValueError("Cannot build a two-class loss with a missing class")
    weights = counts.sum() / (2.0 * counts)
    return nn.CrossEntropyLoss(
        weight=torch.tensor(weights, device=device, dtype=torch.float32),
        label_smoothing=config.label_smoothing,
    )


def _fit_specialist(
    x: np.ndarray,
    y: np.ndarray,
    train_indices: np.ndarray,
    validation_indices: np.ndarray,
    seed: int,
    mode: str,
    config: HardRoutingConfig,
    initial_state: Mapping[str, torch.Tensor] | None,
) -> Tuple[Dict[str, torch.Tensor], Dict[str, object]]:
    seed_everything(seed)
    device = resolve_device(config.device)
    model = ShallowConvNet(channels=8, samples=400).to(device)
    if initial_state is not None:
        model.load_state_dict(initial_state)
    learning_rate = config.global_init_lr if mode == MODEL_GLOBAL_INIT else config.scratch_lr
    max_epochs = config.global_init_epochs if mode == MODEL_GLOBAL_INIT else config.scratch_epochs
    patience = config.global_init_patience if mode == MODEL_GLOBAL_INIT else config.scratch_patience
    batch_size = config.batch_size_fewshot if len(train_indices) + len(validation_indices) <= 40 else config.batch_size_all
    train_loader = make_loader(x, y, train_indices, batch_size, True, stable_seed(seed, "train"), device.type == "cuda")
    validation_loader = make_loader(x, y, validation_indices, batch_size, False, stable_seed(seed, "val"), device.type == "cuda")
    criterion = _criterion(y[train_indices], config, device)
    optimizer = torch.optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=config.weight_decay)
    scheduler = torch.optim.lr_scheduler.LambdaLR(
        optimizer, lambda epoch: _scheduler_lambda(epoch, max_epochs, config.warmup_epochs)
    )
    best_state = None
    best_score = -float("inf")
    best_loss = float("inf")
    best_epoch = 0
    stale = 0
    validation_has_both = len(np.unique(y[validation_indices])) == 2
    train_config = config.training_config()
    for epoch in range(1, max_epochs + 1):
        model.train()
        for raw, target in train_loader:
            raw = raw.to(device, non_blocking=True)
            target = target.to(device, non_blocking=True)
            values = prepare_model_input(raw, "shallow_task", train_config, training=False)
            values = augment_waveforms(values, train_config)
            optimizer.zero_grad(set_to_none=True)
            loss = criterion(model(values), target)
            loss.backward()
            torch.nn.utils.clip_grad_norm_(model.parameters(), config.gradient_clip)
            optimizer.step()
        validation = evaluate_loader(model, validation_loader, "shallow_task", train_config, device, criterion)
        scheduler.step()
        score = float(validation["balanced_accuracy"]) if validation_has_both else -float(validation["loss"])
        loss_value = float(validation["loss"])
        improved = score > best_score + 1e-8 or (
            abs(score - best_score) <= 1e-8 and loss_value < best_loss - 1e-8
        )
        if improved:
            best_state = {key: value.detach().cpu().clone() for key, value in model.state_dict().items()}
            best_score, best_loss, best_epoch, stale = score, loss_value, epoch, 0
        else:
            stale += 1
        if stale >= patience:
            break
    if best_state is None:
        raise RuntimeError("Specialist training produced no checkpoint")
    return best_state, {
        "best_epoch": best_epoch,
        "epochs_ran": epoch,
        "validation_loss": best_loss,
        "validation_has_both_classes": validation_has_both,
        "parameter_count": count_parameters(model),
    }


def _predict_state(
    state: Mapping[str, torch.Tensor],
    x: np.ndarray,
    y: np.ndarray,
    indices: np.ndarray,
    config: HardRoutingConfig,
) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:
    if len(indices) == 0:
        return np.empty(0, dtype=np.int64), np.empty(0, dtype=np.int64), np.empty(0, dtype=np.float32)
    device = resolve_device(config.device)
    model = ShallowConvNet(channels=8, samples=400).to(device)
    model.load_state_dict(state)
    loader = make_loader(x, y, indices, config.batch_size_all, False, stable_seed("predict", *indices.tolist()), device.type == "cuda")
    criterion = nn.CrossEntropyLoss()
    result = evaluate_loader(model, loader, "shallow_task", config.training_config(), device, criterion)
    return result["labels"], result["predictions"], result["probabilities"]


def _metrics(labels: np.ndarray, predictions: np.ndarray, probabilities: np.ndarray) -> Dict[str, object]:
    matrix = confusion_matrix(labels, predictions, labels=(0, 1))
    has_both_classes = len(np.unique(labels)) == 2
    try:
        if not has_both_classes:
            raise ValueError("single-class subset")
        auc = float(roc_auc_score(labels, probabilities))
    except ValueError:
        auc = float("nan")
    return {
        "accuracy": float(accuracy_score(labels, predictions)),
        "balanced_accuracy": float(balanced_accuracy_score(labels, predictions)) if has_both_classes else float("nan"),
        "macro_f1": float(f1_score(labels, predictions, average="macro", zero_division=0)),
        "roc_auc": auc,
        "tn": int(matrix[0, 0]),
        "fp": int(matrix[0, 1]),
        "fn": int(matrix[1, 0]),
        "tp": int(matrix[1, 1]),
    }


def _evaluate_routed(
    states: Mapping[int, Mapping[str, torch.Tensor]],
    global_state: Mapping[str, torch.Tensor],
    fallbacks: Mapping[int, str],
    x: np.ndarray,
    y: np.ndarray,
    test_indices: np.ndarray,
    test_clusters: np.ndarray,
    config: HardRoutingConfig,
) -> Tuple[Dict[str, object], List[Dict[str, object]]]:
    labels = np.empty(len(test_indices), dtype=np.int64)
    predictions = np.empty(len(test_indices), dtype=np.int64)
    probabilities = np.empty(len(test_indices), dtype=np.float32)
    expert_rows: List[Dict[str, object]] = []
    for cluster in range(config.kmeans_clusters):
        positions = np.flatnonzero(test_clusters == cluster)
        routed_indices = test_indices[positions]
        state = global_state if cluster in fallbacks else states[cluster]
        actual, predicted, probability = _predict_state(state, x, y, routed_indices, config)
        labels[positions] = actual
        predictions[positions] = predicted
        probabilities[positions] = probability
        cluster_metrics = _metrics(actual, predicted, probability) if len(actual) else {
            "accuracy": float("nan"), "balanced_accuracy": float("nan"), "macro_f1": float("nan"), "roc_auc": float("nan"),
            "tn": 0, "fp": 0, "fn": 0, "tp": 0,
        }
        expert_rows.append({"cluster": cluster, "fallback_reason": fallbacks.get(cluster, ""), "test_size": len(actual), **cluster_metrics})
    return _metrics(labels, predictions, probabilities), expert_rows


def _save(path: Path, frame: pd.DataFrame) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    temporary = path.with_suffix(path.suffix + ".tmp")
    frame.to_csv(temporary, index=False, encoding="utf-8-sig")
    temporary.replace(path)


def _complete_keys(frame: pd.DataFrame) -> set[Tuple[object, ...]]:
    if frame.empty:
        return set()
    counts = frame.groupby(["split_set", "train_seed", "subject_id", "fold"])["model"].nunique()
    return {tuple(index) for index, count in counts.items() if count == len(MODEL_NAMES)}


def run_evaluation(
    cache: Path,
    output: Path,
    config: HardRoutingConfig | None = None,
    subjects: Sequence[str] = LABELED_SUBJECTS,
    max_new_folds: int | None = None,
) -> Path:
    config = config or HardRoutingConfig()
    output.mkdir(parents=True, exist_ok=True)
    result_path = output / "fold_results.csv"
    diagnostic_path = output / "cluster_diagnostics.csv"
    expert_path = output / "expert_results.csv"
    results = pd.read_csv(result_path) if result_path.exists() else pd.DataFrame()
    diagnostics = pd.read_csv(diagnostic_path) if diagnostic_path.exists() else pd.DataFrame()
    experts = pd.read_csv(expert_path) if expert_path.exists() else pd.DataFrame()
    complete = _complete_keys(results)
    result_rows = results.to_dict("records") if not results.empty else []
    diagnostic_rows = diagnostics.to_dict("records") if not diagnostics.empty else []
    expert_rows = experts.to_dict("records") if not experts.empty else []
    assignments = load_fold_assignments(cache / "fold_assignments.csv")
    split_specs = [("all", config.all_trial_seeds)] + [
        (f"20_seed_{seed}", (config.seed,)) for seed in SUBSAMPLE_SEEDS
    ]
    new_folds = 0
    for split_set, train_seeds in split_specs:
        protocol = "all" if split_set == "all" else "20x20"
        for train_seed in train_seeds:
            for subject in subjects:
                x, y = load_subject(cache, subject)
                features = baseline_features(x)
                for fold in range(1, 6):
                    bundle_key = (split_set, train_seed, subject, fold)
                    if bundle_key in complete:
                        continue
                    split = assignments[(split_set, subject, fold)]
                    fold_seed = stable_seed("kmeans-shallow", split_set, train_seed, subject, fold)
                    print(f"[hard-k2] {protocol} {split_set} seed={train_seed} {subject} fold={fold}", flush=True)
                    router = TrainOnlyKMeansRouter(config).fit(features[split["train"]])
                    routed = {role: router.predict(features[indices]) for role, indices in split.items()}
                    global_seed = stable_seed("outer", "shallow_task", "waveform", split_set, train_seed, subject, fold)
                    global_checkpoint = output / "checkpoints" / subject / f"{split_set}_seed{train_seed}_fold{fold}_global.pt"
                    global_result = train_fold(
                        x, y, split, "shallow_task", "waveform", global_seed,
                        config.training_config(), config.scratch_epochs, config.scratch_patience,
                        evaluate_test=True, checkpoint_path=global_checkpoint,
                    )
                    checkpoint = torch.load(global_checkpoint, map_location="cpu", weights_only=False)
                    global_state = checkpoint["model_state"]
                    common = {
                        "protocol": protocol, "split_set": split_set, "train_seed": train_seed,
                        "subject_id": subject, "fold": fold, "fold_seed": fold_seed,
                    }
                    result_rows.append({**common, **global_result, "model": MODEL_GLOBAL, "routing": "none"})
                    cluster_meta: Dict[int, Dict[str, object]] = {}
                    for cluster in range(config.kmeans_clusters):
                        train_mask = routed["train"] == cluster
                        validation_mask = routed["val"] == cluster
                        test_mask = routed["test"] == cluster
                        train_indices = split["train"][train_mask]
                        validation_indices = split["val"][validation_mask]
                        test_indices = split["test"][test_mask]
                        reason = fallback_reason(y[train_indices], y[validation_indices])
                        train_left, train_right = class_counts(y[split["train"]], routed["train"], cluster)
                        val_left, val_right = class_counts(y[split["val"]], routed["val"], cluster)
                        test_left, test_right = class_counts(y[split["test"]], routed["test"], cluster)
                        cluster_meta[cluster] = {
                            "train_indices": train_indices, "validation_indices": validation_indices,
                            "fallback_reason": reason,
                        }
                        diagnostic_rows.append({
                            **common, "cluster": cluster, "router_fit_samples": router.fitted_samples,
                            "pca_components": router.pca_components, "train_silhouette": router.train_silhouette(features[split["train"]]),
                            "train_left": train_left, "train_right": train_right,
                            "val_left": val_left, "val_right": val_right,
                            "test_left": test_left, "test_right": test_right,
                            "fallback_reason": reason,
                        })
                    for mode in (MODEL_SCRATCH, MODEL_GLOBAL_INIT):
                        states: Dict[int, Mapping[str, torch.Tensor]] = {}
                        fallbacks: Dict[int, str] = {}
                        training_meta: Dict[int, Dict[str, object]] = {}
                        for cluster in range(config.kmeans_clusters):
                            meta = cluster_meta[cluster]
                            reason = str(meta["fallback_reason"])
                            if reason:
                                fallbacks[cluster] = reason
                                training_meta[cluster] = {"best_epoch": 0, "epochs_ran": 0, "validation_loss": float("nan"), "validation_has_both_classes": False, "parameter_count": count_parameters(ShallowConvNet(channels=8, samples=400))}
                                continue
                            expert_seed = stable_seed("expert", mode, split_set, train_seed, subject, fold, cluster)
                            initial_state = global_state if mode == MODEL_GLOBAL_INIT else None
                            state, meta_result = _fit_specialist(
                                x, y, meta["train_indices"], meta["validation_indices"], expert_seed,
                                mode, config, initial_state,
                            )
                            states[cluster] = state
                            training_meta[cluster] = meta_result
                        routed_metrics, per_expert = _evaluate_routed(
                            states, global_state, fallbacks, x, y, split["test"], routed["test"], config
                        )
                        result_rows.append({
                            **common, "model": mode, "routing": "baseline_train_only_kmeans_k2",
                            "representation": "task_waveform", "train_size": len(split["train"]),
                            "val_size": len(split["val"]), "test_size": len(split["test"]),
                            "fallback_clusters": len(fallbacks), "parameter_count_per_expert": count_parameters(ShallowConvNet(channels=8, samples=400)),
                            **routed_metrics,
                        })
                        for per_cluster in per_expert:
                            cluster = int(per_cluster["cluster"])
                            expert_rows.append({
                                **common, "model": mode, **training_meta[cluster], **per_cluster,
                            })
                    _save(result_path, pd.DataFrame(result_rows))
                    _save(diagnostic_path, pd.DataFrame(diagnostic_rows))
                    _save(expert_path, pd.DataFrame(expert_rows))
                    complete.add(bundle_key)
                    new_folds += 1
                    if max_new_folds is not None and new_folds >= max_new_folds:
                        return result_path
    summarize(output)
    (output / "run_config.json").write_text(json.dumps(asdict(config), indent=2), encoding="utf-8")
    return result_path


def _bootstrap_ci(values: np.ndarray, seed: int, draws: int = 10000) -> Tuple[float, float]:
    rng = np.random.default_rng(seed)
    sampled = rng.choice(values, size=(draws, len(values)), replace=True).mean(axis=1)
    return tuple(np.quantile(sampled, (0.025, 0.975)).tolist())


def summarize(output: Path) -> None:
    folds = pd.read_csv(output / "fold_results.csv")
    subject = folds.groupby(["protocol", "model", "subject_id"], as_index=False).agg(
        mean_accuracy=("accuracy", "mean"), std_accuracy=("accuracy", "std"),
        mean_balanced_accuracy=("balanced_accuracy", "mean"), std_balanced_accuracy=("balanced_accuracy", "std"),
        mean_macro_f1=("macro_f1", "mean"), fallback_cluster_rate=("fallback_clusters", "mean"),
    )
    group = subject.groupby(["protocol", "model"], as_index=False).agg(
        subjects=("subject_id", "nunique"), mean_accuracy=("mean_accuracy", "mean"),
        std_subject_accuracy=("mean_accuracy", "std"), mean_balanced_accuracy=("mean_balanced_accuracy", "mean"),
        std_subject_balanced_accuracy=("mean_balanced_accuracy", "std"), mean_macro_f1=("mean_macro_f1", "mean"),
        mean_fallback_clusters=("fallback_cluster_rate", "mean"),
    )
    comparisons = []
    for protocol in ("all", "20x20"):
        pivot = subject[subject.protocol == protocol].pivot(index="subject_id", columns="model", values="mean_accuracy")
        for model in (MODEL_SCRATCH, MODEL_GLOBAL_INIT):
            differences = (pivot[model] - pivot[MODEL_GLOBAL]).dropna().to_numpy()
            low, high = _bootstrap_ci(differences, stable_seed("bootstrap", protocol, model))
            try:
                p_value = float(wilcoxon(differences, alternative="two-sided", method="auto").pvalue)
            except ValueError:
                p_value = float("nan")
            comparisons.append({
                "protocol": protocol, "comparison": f"{model} - {MODEL_GLOBAL}", "subjects": len(differences),
                "mean_difference_pp": 100 * float(differences.mean()), "bootstrap_ci_low_pp": 100 * low,
                "bootstrap_ci_high_pp": 100 * high, "wilcoxon_p": p_value,
                "improved_subjects": int((differences > 0).sum()),
            })
    _save(output / "subject_summary.csv", subject)
    _save(output / "group_summary.csv", group)
    _save(output / "paired_comparisons.csv", pd.DataFrame(comparisons))


def smoke_test(cache: Path, device: str = "cpu") -> Dict[str, object]:
    config = HardRoutingConfig(device=device, scratch_epochs=2, scratch_patience=2, global_init_epochs=2, global_init_patience=2)
    x, y = load_subject(cache, "sub1")
    assignments = load_fold_assignments(cache / "fold_assignments.csv")
    split = assignments[(f"20_seed_{SUBSAMPLE_SEEDS[0]}", "sub1", 1)]
    features = baseline_features(x)
    router = TrainOnlyKMeansRouter(config).fit(features[split["train"]])
    train_clusters = router.predict(features[split["train"]])
    validation_clusters = router.predict(features[split["val"]])
    model = ShallowConvNet(channels=8, samples=400)
    dummy = torch.randn(2, 1, 8, 400)
    logits = model(dummy)
    if tuple(logits.shape) != (2, 2):
        raise RuntimeError("Invalid ShallowConvNet output")
    reasons = []
    for cluster in range(2):
        train_indices = split["train"][train_clusters == cluster]
        val_indices = split["val"][validation_clusters == cluster]
        reasons.append(fallback_reason(y[train_indices], y[val_indices]))
    checks = {
        "subjects": len(LABELED_SUBJECTS), "feature_shape": list(features.shape),
        "router_fit_samples": router.fitted_samples, "pca_components": router.pca_components,
        "train_cluster_counts": np.bincount(train_clusters, minlength=2).tolist(),
        "fallback_reasons": reasons, "shallow_parameters": count_parameters(model),
        "output_shape": list(logits.shape), "device": device,
    }
    return checks
