#!/usr/bin/env python3
"""Search single-bit flips in a CIFAR-10 classifier's final layer.

The search is exhaustive over every stored bit in the 10 x 64 final-layer
weight matrix for three storage formats: IEEE-like float32, float16, and
per-tensor symmetric int8. The convolutional backbone remains float32 in all
three comparisons so the experiment isolates final-layer storage.
"""

from __future__ import annotations

import argparse
import csv
import hashlib
import json
import math
import platform
import sys
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Iterable

import matplotlib.pyplot as plt
import numpy as np
import torch
import torchvision
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
from torchvision.datasets.utils import download_url


MODEL_REPOSITORY = "https://github.com/chenyaofo/pytorch-cifar-models"
MODEL_SOURCE_COMMIT = "786c16252c0fc58ee9adac063f8337cc4a7a497a"
MODEL_HUB_SPEC = f"chenyaofo/pytorch-cifar-models:{MODEL_SOURCE_COMMIT}"
MODEL_NAME = "cifar10_resnet20"
CHECKPOINT_FILENAME = "cifar10_resnet20-4118986f.pt"
CHECKPOINT_SHA256 = "4118986f0df73003d572b0e397f0ac7b3f60af1f31aff3d2da164536e36f6ec8"
CIFAR10_ARCHIVE_MD5 = "c58f30108f718f92721af3b95e74349a"
CIFAR10_MIRROR_URL = "https://dataset.bj.bcebos.com/cifar/cifar-10-python.tar.gz"
CIFAR10_CLASSES = [
    "airplane",
    "automobile",
    "bird",
    "cat",
    "deer",
    "dog",
    "frog",
    "horse",
    "ship",
    "truck",
]
ANIMAL_CLASS_INDICES = (2, 3, 4, 5, 6, 7)
NORMALIZE_MEAN = (0.4914, 0.4822, 0.4465)
NORMALIZE_STD = (0.2023, 0.1994, 0.2010)


@dataclass(frozen=True)
class SearchResult:
    storage_format: str
    target_index: int
    feature_index: int
    flat_index: int
    bit_index: int
    bit_kind: str
    old_bits: int
    new_bits: int
    old_value: float
    new_value: float
    calibration_accuracy: float
    calibration_target_share: float

    @property
    def target_class(self) -> str:
        return CIFAR10_CLASSES[self.target_index]

    def ranking_key(self) -> tuple[float, float, int, int]:
        return (
            self.calibration_target_share,
            -self.calibration_accuracy,
            -self.flat_index,
            -self.bit_index,
        )


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--device",
        default="auto",
        choices=("auto", "cpu", "cuda"),
        help="Device used only for ResNet feature extraction.",
    )
    parser.add_argument("--batch-size", type=int, default=512)
    parser.add_argument("--calibration-per-class", type=int, default=100)
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=Path(__file__).resolve().parent / "outputs",
    )
    parser.add_argument(
        "--data-dir",
        type=Path,
        default=Path(__file__).resolve().parent / ".cache" / "data",
    )
    return parser.parse_args()


def sha256_file(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for block in iter(lambda: handle.read(1024 * 1024), b""):
            digest.update(block)
    return digest.hexdigest()


def md5_file(path: Path) -> str:
    digest = hashlib.md5(usedforsecurity=False)
    with path.open("rb") as handle:
        for block in iter(lambda: handle.read(1024 * 1024), b""):
            digest.update(block)
    return digest.hexdigest()


def choose_device(requested: str) -> torch.device:
    if requested == "cuda":
        if not torch.cuda.is_available():
            raise RuntimeError("--device cuda requested, but CUDA is unavailable")
        return torch.device("cuda")
    if requested == "cpu":
        return torch.device("cpu")
    return torch.device("cuda" if torch.cuda.is_available() else "cpu")


def load_model(device: torch.device) -> tuple[torch.nn.Module, Path]:
    model = torch.hub.load(
        MODEL_HUB_SPEC,
        MODEL_NAME,
        pretrained=True,
        trust_repo=True,
        skip_validation=True,
        verbose=True,
    )
    checkpoint = Path(torch.hub.get_dir()) / "checkpoints" / CHECKPOINT_FILENAME
    if not checkpoint.exists():
        raise FileNotFoundError(f"PyTorch Hub did not leave the checkpoint at {checkpoint}")
    actual_hash = sha256_file(checkpoint)
    if actual_hash != CHECKPOINT_SHA256:
        raise RuntimeError(
            f"checkpoint SHA-256 mismatch: expected {CHECKPOINT_SHA256}, got {actual_hash}"
        )
    model.eval().to(device)
    return model, checkpoint


def load_test_data(data_dir: Path, batch_size: int) -> tuple[datasets.CIFAR10, DataLoader]:
    data_dir.mkdir(parents=True, exist_ok=True)
    download_url(
        CIFAR10_MIRROR_URL,
        str(data_dir),
        filename="cifar-10-python.tar.gz",
        md5=CIFAR10_ARCHIVE_MD5,
    )
    transform = transforms.Compose(
        [
            transforms.ToTensor(),
            transforms.Normalize(NORMALIZE_MEAN, NORMALIZE_STD),
        ]
    )
    dataset = datasets.CIFAR10(root=data_dir, train=False, download=True, transform=transform)
    loader = DataLoader(
        dataset,
        batch_size=batch_size,
        shuffle=False,
        num_workers=0,
        pin_memory=torch.cuda.is_available(),
    )
    return dataset, loader


def extract_features_and_logits(
    model: torch.nn.Module, loader: DataLoader, device: torch.device
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
    feature_batches: list[torch.Tensor] = []
    logit_batches: list[torch.Tensor] = []
    label_batches: list[torch.Tensor] = []

    def capture_fc_input(_module: torch.nn.Module, inputs: tuple[torch.Tensor, ...]) -> None:
        feature_batches.append(inputs[0].detach().cpu())

    hook = model.fc.register_forward_pre_hook(capture_fc_input)
    try:
        with torch.inference_mode():
            for images, labels in loader:
                logits = model(images.to(device, non_blocking=True))
                logit_batches.append(logits.detach().cpu())
                label_batches.append(labels)
    finally:
        hook.remove()

    features = torch.cat(feature_batches).numpy().astype(np.float32, copy=False)
    logits = torch.cat(logit_batches).numpy().astype(np.float32, copy=False)
    labels = torch.cat(label_batches).numpy().astype(np.int64, copy=False)
    return features, logits, labels


def evaluate_model_logits(
    model: torch.nn.Module, loader: DataLoader, device: torch.device
) -> np.ndarray:
    batches: list[torch.Tensor] = []
    with torch.inference_mode():
        for images, _labels in loader:
            batches.append(model(images.to(device, non_blocking=True)).detach().cpu())
    return torch.cat(batches).numpy().astype(np.float32, copy=False)


def stratified_calibration_indices(labels: np.ndarray, per_class: int) -> np.ndarray:
    selected: list[int] = []
    for class_index in range(len(CIFAR10_CLASSES)):
        class_rows = np.flatnonzero(labels == class_index)
        if len(class_rows) < per_class:
            raise ValueError(f"class {class_index} has only {len(class_rows)} rows")
        selected.extend(class_rows[:per_class].tolist())
    return np.array(sorted(selected), dtype=np.int64)


def metrics(logits: np.ndarray, labels: np.ndarray) -> dict[str, object]:
    predictions = logits.argmax(axis=1)
    distribution = np.bincount(predictions, minlength=len(CIFAR10_CLASSES))
    per_class_accuracy: list[float] = []
    for class_index in range(len(CIFAR10_CLASSES)):
        mask = labels == class_index
        per_class_accuracy.append(float(np.mean(predictions[mask] == labels[mask])))
    return {
        "accuracy": float(np.mean(predictions == labels)),
        "prediction_counts": distribution.astype(int).tolist(),
        "prediction_shares": (distribution / len(labels)).astype(float).tolist(),
        "per_class_accuracy": per_class_accuracy,
        "predictions": predictions,
    }


def bit_kind(storage_format: str, bit_index: int) -> str:
    if storage_format == "float32":
        if bit_index == 31:
            return "sign"
        if bit_index >= 23:
            return "exponent"
        return "mantissa"
    if storage_format == "float16":
        if bit_index == 15:
            return "sign"
        if bit_index >= 10:
            return "exponent"
        return "mantissa"
    return "integer-storage"


def best_other_classes(logits: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
    class_count = logits.shape[1]
    other_values = np.empty_like(logits)
    other_indices = np.empty(logits.shape, dtype=np.int64)
    for target in range(class_count):
        masked = logits.copy()
        masked[:, target] = -np.inf
        indices = masked.argmax(axis=1)
        other_indices[:, target] = indices
        other_values[:, target] = masked[np.arange(len(masked)), indices]
    return other_values, other_indices


def score_candidate(
    baseline_logits: np.ndarray,
    features: np.ndarray,
    labels: np.ndarray,
    other_values: np.ndarray,
    other_indices: np.ndarray,
    target_index: int,
    feature_index: int,
    delta: float,
) -> tuple[float, float]:
    mutated_target = (
        baseline_logits[:, target_index].astype(np.float64)
        + features[:, feature_index].astype(np.float64) * float(delta)
    )
    target_wins = mutated_target > other_values[:, target_index].astype(np.float64)
    predictions = np.where(target_wins, target_index, other_indices[:, target_index])
    return float(np.mean(predictions == labels)), float(np.mean(predictions == target_index))


def float32_flip(value: np.float32, bit_index: int) -> tuple[int, int, float]:
    old_bits = int(np.array([value], dtype="<f4").view("<u4")[0])
    new_bits = old_bits ^ (1 << bit_index)
    new_value = float(np.array([new_bits], dtype="<u4").view("<f4")[0])
    return old_bits, new_bits, new_value


def float16_flip(value: np.float16, bit_index: int) -> tuple[int, int, float]:
    old_bits = int(np.array([value], dtype="<f2").view("<u2")[0])
    new_bits = old_bits ^ (1 << bit_index)
    new_value = float(np.array([new_bits], dtype="<u2").view("<f2")[0])
    return old_bits, new_bits, new_value


def int8_flip(value: np.int8, bit_index: int) -> tuple[int, int, int]:
    old_bits = int(np.array([value], dtype=np.int8).view(np.uint8)[0])
    new_bits = old_bits ^ (1 << bit_index)
    new_value = int(np.array([new_bits], dtype=np.uint8).view(np.int8)[0])
    return old_bits, new_bits, new_value


def search_storage_format(
    storage_format: str,
    weights: np.ndarray,
    bias: np.ndarray,
    features: np.ndarray,
    labels: np.ndarray,
    csv_writer: csv.DictWriter,
) -> tuple[SearchResult, dict[str, dict[str, float | int]], np.ndarray, np.ndarray, float | None]:
    if storage_format == "float32":
        stored_weights = weights.astype(np.float32, copy=True)
        effective_weights = stored_weights.astype(np.float32)
        effective_bias = bias.astype(np.float32)
        bits_per_weight = 32
        scale = None
    elif storage_format == "float16":
        stored_weights = weights.astype(np.float16)
        effective_weights = stored_weights.astype(np.float32)
        effective_bias = bias.astype(np.float16).astype(np.float32)
        bits_per_weight = 16
        scale = None
    elif storage_format == "int8":
        scale = float(np.max(np.abs(weights)) / 127.0)
        stored_weights = np.clip(np.rint(weights / scale), -127, 127).astype(np.int8)
        effective_weights = stored_weights.astype(np.float32) * scale
        effective_bias = bias.astype(np.float32)
        bits_per_weight = 8
    else:
        raise ValueError(storage_format)

    baseline_logits = features @ effective_weights.T + effective_bias
    baseline = metrics(baseline_logits, labels)
    other_values, other_indices = best_other_classes(baseline_logits)
    best: SearchResult | None = None
    field_summary: dict[str, dict[str, float | int]] = {}

    for target_index in range(stored_weights.shape[0]):
        for feature_index in range(stored_weights.shape[1]):
            flat_index = target_index * stored_weights.shape[1] + feature_index
            old_stored_value = stored_weights[target_index, feature_index]
            old_effective_value = float(effective_weights[target_index, feature_index])
            for bit_index in range(bits_per_weight):
                kind = bit_kind(storage_format, bit_index)
                if storage_format == "float32":
                    old_bits, new_bits, new_effective_value = float32_flip(
                        np.float32(old_stored_value), bit_index
                    )
                elif storage_format == "float16":
                    old_bits, new_bits, new_effective_value = float16_flip(
                        np.float16(old_stored_value), bit_index
                    )
                else:
                    old_bits, new_bits, new_stored_value = int8_flip(
                        np.int8(old_stored_value), bit_index
                    )
                    new_effective_value = float(new_stored_value * scale)

                is_finite = math.isfinite(new_effective_value)
                if is_finite:
                    accuracy, target_share = score_candidate(
                        baseline_logits,
                        features,
                        labels,
                        other_values,
                        other_indices,
                        target_index,
                        feature_index,
                        new_effective_value - old_effective_value,
                    )
                else:
                    accuracy, target_share = math.nan, math.nan

                csv_writer.writerow(
                    {
                        "storage_format": storage_format,
                        "target_class": CIFAR10_CLASSES[target_index],
                        "target_index": target_index,
                        "feature_index": feature_index,
                        "flat_index": flat_index,
                        "bit_index": bit_index,
                        "bit_kind": kind,
                        "old_bits_hex": f"0x{old_bits:0{bits_per_weight // 4}x}",
                        "new_bits_hex": f"0x{new_bits:0{bits_per_weight // 4}x}",
                        "old_effective_value": f"{old_effective_value:.17g}",
                        "new_effective_value": f"{new_effective_value:.17g}",
                        "finite": str(is_finite).lower(),
                        "calibration_accuracy": "" if not is_finite else f"{accuracy:.12f}",
                        "calibration_target_share": "" if not is_finite else f"{target_share:.12f}",
                    }
                )

                summary = field_summary.setdefault(
                    kind,
                    {
                        "candidates": 0,
                        "finite_candidates": 0,
                        "max_animal_target_share": 0.0,
                        "min_accuracy": 1.0,
                    },
                )
                summary["candidates"] = int(summary["candidates"]) + 1
                if is_finite:
                    summary["finite_candidates"] = int(summary["finite_candidates"]) + 1
                    summary["min_accuracy"] = min(float(summary["min_accuracy"]), accuracy)
                    if target_index in ANIMAL_CLASS_INDICES:
                        summary["max_animal_target_share"] = max(
                            float(summary["max_animal_target_share"]), target_share
                        )
                        candidate = SearchResult(
                            storage_format=storage_format,
                            target_index=target_index,
                            feature_index=feature_index,
                            flat_index=flat_index,
                            bit_index=bit_index,
                            bit_kind=kind,
                            old_bits=old_bits,
                            new_bits=new_bits,
                            old_value=old_effective_value,
                            new_value=new_effective_value,
                            calibration_accuracy=accuracy,
                            calibration_target_share=target_share,
                        )
                        if best is None or candidate.ranking_key() > best.ranking_key():
                            best = candidate

    if best is None:
        raise RuntimeError(f"no finite animal-target candidate found for {storage_format}")
    return (
        best,
        field_summary,
        baseline_logits,
        baseline["predictions"],
        scale,
    )


def apply_winner(
    winner: SearchResult,
    weights: np.ndarray,
    bias: np.ndarray,
    features: np.ndarray,
    int8_scale: float | None,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
    if winner.storage_format == "float32":
        effective_weights = weights.astype(np.float32, copy=True)
        effective_bias = bias.astype(np.float32, copy=True)
    elif winner.storage_format == "float16":
        effective_weights = weights.astype(np.float16).astype(np.float32)
        effective_bias = bias.astype(np.float16).astype(np.float32)
    elif winner.storage_format == "int8":
        if int8_scale is None:
            raise ValueError("int8 scale missing")
        quantized = np.clip(np.rint(weights / int8_scale), -127, 127).astype(np.int8)
        effective_weights = quantized.astype(np.float32) * int8_scale
        effective_bias = bias.astype(np.float32, copy=True)
    else:
        raise ValueError(winner.storage_format)

    original = effective_weights.copy()
    effective_weights[winner.target_index, winner.feature_index] = winner.new_value
    with np.errstate(over="ignore", invalid="ignore"):
        mutated_logits = features @ effective_weights.T + effective_bias
    return mutated_logits, original, effective_weights


def format_winner(winner: SearchResult, bit_width: int) -> dict[str, object]:
    return {
        "storage_format": winner.storage_format,
        "parameter": "fc.weight",
        "target_class": winner.target_class,
        "target_index": winner.target_index,
        "feature_index": winner.feature_index,
        "flat_index": winner.flat_index,
        "bit_index_from_lsb": winner.bit_index,
        "bit_kind": winner.bit_kind,
        "old_bits_hex": f"0x{winner.old_bits:0{bit_width // 4}x}",
        "new_bits_hex": f"0x{winner.new_bits:0{bit_width // 4}x}",
        "old_effective_value": winner.old_value,
        "new_effective_value": winner.new_value,
        "calibration_accuracy": winner.calibration_accuracy,
        "calibration_target_share": winner.calibration_target_share,
    }


def write_distribution_csv(
    path: Path, scenarios: dict[str, dict[str, object]], total: int
) -> None:
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(
            handle,
            fieldnames=("scenario", "class", "prediction_count", "prediction_share"),
            lineterminator="\n",
        )
        writer.writeheader()
        for scenario, scenario_metrics in scenarios.items():
            counts = scenario_metrics["prediction_counts"]
            for class_name, count in zip(CIFAR10_CLASSES, counts):
                writer.writerow(
                    {
                        "scenario": scenario,
                        "class": class_name,
                        "prediction_count": count,
                        "prediction_share": f"{count / total:.12f}",
                    }
                )


def write_per_class_csv(path: Path, scenarios: dict[str, dict[str, object]]) -> None:
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(
            handle,
            fieldnames=("scenario", "class", "accuracy"),
            lineterminator="\n",
        )
        writer.writeheader()
        for scenario, scenario_metrics in scenarios.items():
            for class_name, accuracy in zip(
                CIFAR10_CLASSES, scenario_metrics["per_class_accuracy"]
            ):
                writer.writerow(
                    {"scenario": scenario, "class": class_name, "accuracy": f"{accuracy:.12f}"}
                )


def write_winner_predictions_csv(
    path: Path,
    labels: np.ndarray,
    baseline_logits: np.ndarray,
    mutated_logits: np.ndarray,
    feature_activations: np.ndarray,
    target_index: int,
) -> None:
    baseline_predictions = baseline_logits.argmax(axis=1)
    mutated_predictions = mutated_logits.argmax(axis=1)
    with path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(
            handle,
            fieldnames=(
                "test_index",
                "true_class",
                "baseline_prediction",
                "winner_feature_activation",
                "original_target_logit",
                "mutated_target_logit",
                "mutated_prediction",
                "prediction_changed",
            ),
            lineterminator="\n",
        )
        writer.writeheader()
        for index in range(len(labels)):
            writer.writerow(
                {
                    "test_index": index,
                    "true_class": CIFAR10_CLASSES[int(labels[index])],
                    "baseline_prediction": CIFAR10_CLASSES[int(baseline_predictions[index])],
                    "winner_feature_activation": f"{float(feature_activations[index]):.12g}",
                    "original_target_logit": f"{float(baseline_logits[index, target_index]):.12g}",
                    "mutated_target_logit": f"{float(mutated_logits[index, target_index]):.12g}",
                    "mutated_prediction": CIFAR10_CLASSES[int(mutated_predictions[index])],
                    "prediction_changed": str(
                        bool(baseline_predictions[index] != mutated_predictions[index])
                    ).lower(),
                }
            )


def summarize_animal_targets(candidate_path: Path, output_path: Path) -> list[dict[str, object]]:
    summary: dict[tuple[str, str], dict[str, object]] = {}
    with candidate_path.open(newline="", encoding="utf-8") as handle:
        for row in csv.DictReader(handle):
            if row["target_class"] not in {
                CIFAR10_CLASSES[index] for index in ANIMAL_CLASS_INDICES
            }:
                continue
            key = (row["storage_format"], row["target_class"])
            item = summary.setdefault(
                key,
                {
                    "storage_format": row["storage_format"],
                    "target_class": row["target_class"],
                    "finite_candidates": 0,
                    "max_calibration_target_share": 0.0,
                    "perfect_calibration_candidates": 0,
                },
            )
            if row["finite"] != "true":
                continue
            share = float(row["calibration_target_share"])
            item["finite_candidates"] = int(item["finite_candidates"]) + 1
            item["max_calibration_target_share"] = max(
                float(item["max_calibration_target_share"]), share
            )
            if share == 1.0:
                item["perfect_calibration_candidates"] = (
                    int(item["perfect_calibration_candidates"]) + 1
                )

    rows = [summary[key] for key in sorted(summary)]
    with output_path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(
            handle,
            fieldnames=(
                "storage_format",
                "target_class",
                "finite_candidates",
                "max_calibration_target_share",
                "perfect_calibration_candidates",
            ),
            lineterminator="\n",
        )
        writer.writeheader()
        for row in rows:
            writer.writerow(
                {
                    **row,
                    "max_calibration_target_share": (
                        f"{float(row['max_calibration_target_share']):.12f}"
                    ),
                }
            )
    return rows


def draw_evidence_chart(
    path: Path,
    baseline_metrics: dict[str, object],
    mutated_metrics: dict[str, object],
    favorite_class: str,
) -> None:
    plt.rcParams.update(
        {
            "figure.facecolor": "white",
            "axes.facecolor": "white",
            "axes.edgecolor": "black",
            "axes.linewidth": 2.5,
            "font.family": "DejaVu Sans",
            "font.size": 12,
            "path.sketch": (1.0, 120.0, 2.0),
        }
    )
    figure, axes = plt.subplots(1, 2, figsize=(16, 7), constrained_layout=True)
    positions = np.arange(len(CIFAR10_CLASSES))
    width = 0.38
    baseline_shares = np.array(baseline_metrics["prediction_shares"]) * 100
    mutated_shares = np.array(mutated_metrics["prediction_shares"]) * 100
    axes[0].bar(
        positions - width / 2,
        baseline_shares,
        width,
        color="#62a7ff",
        edgecolor="black",
        linewidth=1.5,
        label="original",
    )
    axes[0].bar(
        positions + width / 2,
        mutated_shares,
        width,
        color="#ff5a78",
        edgecolor="black",
        linewidth=1.5,
        label="one bit flipped",
    )
    axes[0].set_xticks(positions, CIFAR10_CLASSES, rotation=35, ha="right")
    axes[0].set_ylabel("share of 10,000 predictions (%)")
    axes[0].set_title(f"THE MODEL'S NEW FAVORITE: {favorite_class.upper()}")
    axes[0].legend(frameon=True, edgecolor="black")
    axes[0].grid(axis="y", color="#cccccc", linewidth=0.8)

    accuracies = [
        float(baseline_metrics["accuracy"]) * 100,
        float(mutated_metrics["accuracy"]) * 100,
    ]
    bars = axes[1].bar(
        ["original", "one bit flipped"],
        accuracies,
        color=["#62a7ff", "#ff5a78"],
        edgecolor="black",
        linewidth=2,
    )
    axes[1].set_ylim(0, 100)
    axes[1].set_ylabel("top-1 accuracy (%)")
    axes[1].set_title("ACCURACY LEFT THE BUILDING")
    axes[1].grid(axis="y", color="#cccccc", linewidth=0.8)
    for bar, value in zip(bars, accuracies):
        axes[1].text(
            bar.get_x() + bar.get_width() / 2,
            value + 2,
            f"{value:.2f}%",
            ha="center",
            va="bottom",
            fontsize=16,
            weight="bold",
        )
    figure.suptitle(
        "EXACT TEST-SET OUTPUT FROM CIFAR-10 RESNET-20 (DRAWN BADLY ON PURPOSE)",
        fontsize=17,
        weight="bold",
    )
    figure.savefig(path, dpi=150)
    plt.close(figure)


def json_ready_metrics(value: dict[str, object]) -> dict[str, object]:
    return {key: item for key, item in value.items() if key != "predictions"}


def main() -> int:
    args = parse_args()
    if args.calibration_per_class < 1 or args.calibration_per_class > 1000:
        raise ValueError("--calibration-per-class must be between 1 and 1000")
    args.output_dir.mkdir(parents=True, exist_ok=True)

    torch.manual_seed(0)
    np.random.seed(0)
    torch.use_deterministic_algorithms(True)
    if torch.cuda.is_available():
        torch.backends.cuda.matmul.allow_tf32 = False
        torch.backends.cudnn.allow_tf32 = False
        torch.backends.cudnn.benchmark = False

    started = time.time()
    device = choose_device(args.device)
    print(f"Loading pinned {MODEL_NAME} on {device}...")
    model, checkpoint_path = load_model(device)
    dataset, loader = load_test_data(args.data_dir, args.batch_size)
    print(f"Extracting final-layer features for {len(dataset)} CIFAR-10 test images...")
    full_features, model_logits, labels = extract_features_and_logits(model, loader, device)

    weights = model.fc.weight.detach().cpu().numpy().astype(np.float32, copy=True)
    bias = model.fc.bias.detach().cpu().numpy().astype(np.float32, copy=True)
    reconstructed_logits = full_features @ weights.T + bias
    if not np.allclose(model_logits, reconstructed_logits, rtol=1e-5, atol=1e-5):
        max_difference = float(np.max(np.abs(model_logits - reconstructed_logits)))
        raise AssertionError(f"captured features do not reconstruct logits; max diff {max_difference}")

    original_weight_bytes = weights.tobytes()
    baseline_full = metrics(model_logits, labels)
    calibration_indices = stratified_calibration_indices(labels, args.calibration_per_class)
    calibration_features = full_features[calibration_indices]
    calibration_labels = labels[calibration_indices]

    candidate_csv_path = args.output_dir / "candidate-search.csv"
    csv_fields = (
        "storage_format",
        "target_class",
        "target_index",
        "feature_index",
        "flat_index",
        "bit_index",
        "bit_kind",
        "old_bits_hex",
        "new_bits_hex",
        "old_effective_value",
        "new_effective_value",
        "finite",
        "calibration_accuracy",
        "calibration_target_share",
    )
    winners: dict[str, SearchResult] = {}
    summaries: dict[str, dict[str, dict[str, float | int]]] = {}
    format_baselines: dict[str, dict[str, object]] = {}
    int8_scale: float | None = None

    with candidate_csv_path.open("w", newline="", encoding="utf-8") as handle:
        writer = csv.DictWriter(
            handle,
            fieldnames=csv_fields,
            lineterminator="\n",
        )
        writer.writeheader()
        for storage_format in ("float32", "float16", "int8"):
            print(f"Searching every final-layer {storage_format} bit on the calibration set...")
            winner, summary, calibration_logits, _, scale = search_storage_format(
                storage_format,
                weights,
                bias,
                calibration_features,
                calibration_labels,
                writer,
            )
            winners[storage_format] = winner
            summaries[storage_format] = summary
            format_baselines[storage_format] = json_ready_metrics(
                metrics(calibration_logits, calibration_labels)
            )
            if storage_format == "int8":
                int8_scale = scale
            print(
                f"  best animal target: {winner.target_class}, bit {winner.bit_index} "
                f"({winner.bit_kind}), share {winner.calibration_target_share:.4%}"
            )

    animal_summary_path = args.output_dir / "animal-target-summary.csv"
    animal_target_summaries = summarize_animal_targets(
        candidate_csv_path, animal_summary_path
    )

    full_scenarios: dict[str, dict[str, object]] = {"float32_original": baseline_full}
    full_winner_details: dict[str, dict[str, object]] = {}
    for storage_format, winner in winners.items():
        mutated_logits, original_effective, mutated_effective = apply_winner(
            winner, weights, bias, full_features, int8_scale
        )
        width = {"float32": 32, "float16": 16, "int8": 8}[storage_format]
        full_metrics = metrics(mutated_logits, labels)
        scenario = f"{storage_format}_winner"
        full_scenarios[scenario] = full_metrics
        changed = np.flatnonzero(original_effective.view(np.uint8) != mutated_effective.view(np.uint8))
        full_winner_details[storage_format] = {
            **format_winner(winner, width),
            "full_test_accuracy": full_metrics["accuracy"],
            "full_test_target_share": full_metrics["prediction_shares"][winner.target_index],
            "full_test_prediction_count": full_metrics["prediction_counts"][winner.target_index],
            "effective_weight_elements_changed": int(
                np.count_nonzero(original_effective != mutated_effective)
            ),
            "effective_weight_bytes_changed_after_conversion": int(len(changed)),
            "storage_weight_bytes_changed": 1,
            "nonfinite_output_logits": int(np.count_nonzero(~np.isfinite(mutated_logits))),
        }

    float32_winner = winners["float32"]
    with torch.no_grad():
        parameter = model.fc.weight[float32_winner.target_index, float32_winner.feature_index]
        original_parameter_value = parameter.detach().cpu().numpy().astype("<f4")
        original_parameter_bits = int(original_parameter_value.view("<u4"))
        if original_parameter_bits != float32_winner.old_bits:
            raise AssertionError("selected source bits do not match the live model parameter")
        parameter.copy_(torch.tensor(float32_winner.new_value, device=device, dtype=torch.float32))
        mutated_parameter_bits = int(
            parameter.detach().cpu().numpy().astype("<f4").view("<u4")
        )
        if mutated_parameter_bits != float32_winner.new_bits:
            raise AssertionError("live model parameter did not receive the selected bit pattern")

    actual_mutated_logits = evaluate_model_logits(model, loader, device)
    actual_mutated_metrics = metrics(actual_mutated_logits, labels)
    analytical_predictions = full_scenarios["float32_winner"]["predictions"]
    actual_predictions = actual_mutated_metrics["predictions"]
    if not np.array_equal(actual_predictions, analytical_predictions):
        disagreements = int(np.count_nonzero(actual_predictions != analytical_predictions))
        raise AssertionError(f"live model and analytical winner disagree on {disagreements} predictions")
    full_scenarios["float32_winner"] = actual_mutated_metrics
    full_winner_details["float32"].update(
        {
            "full_test_accuracy": actual_mutated_metrics["accuracy"],
            "full_test_target_share": actual_mutated_metrics["prediction_shares"][
                float32_winner.target_index
            ],
            "full_test_prediction_count": actual_mutated_metrics["prediction_counts"][
                float32_winner.target_index
            ],
            "nonfinite_output_logits": int(
                np.count_nonzero(~np.isfinite(actual_mutated_logits))
            ),
        }
    )
    with torch.no_grad():
        parameter.copy_(
            torch.tensor(float32_winner.old_value, device=device, dtype=torch.float32)
        )
        restored_parameter_bits = int(
            parameter.detach().cpu().numpy().astype("<f4").view("<u4")
        )
    if restored_parameter_bits != float32_winner.old_bits:
        raise AssertionError("flipping the selected parameter back did not restore its source bits")

    float32_mutated = full_scenarios["float32_winner"]
    winner_feature_activations = full_features[:, float32_winner.feature_index]
    favorite_mask = actual_predictions == float32_winner.target_index
    positive_feature_mask = winner_feature_activations > 0
    favorite_matches_positive_feature = bool(
        np.array_equal(favorite_mask, positive_feature_mask)
    )
    if not favorite_matches_positive_feature:
        raise AssertionError("favorite predictions do not match positive winning-feature activations")
    full_winner_details["float32"]["winner_feature_activation"] = {
        "feature_index": float32_winner.feature_index,
        "positive_count": int(np.count_nonzero(positive_feature_mask)),
        "zero_count": int(np.count_nonzero(winner_feature_activations == 0)),
        "negative_count": int(np.count_nonzero(winner_feature_activations < 0)),
        "minimum": float(np.min(winner_feature_activations)),
        "median": float(np.median(winner_feature_activations)),
        "maximum": float(np.max(winner_feature_activations)),
        "favorite_predictions_exactly_match_positive_activations": True,
    }
    premise_passed = bool(
        float32_mutated["prediction_shares"][float32_winner.target_index] >= 0.90
        and float32_mutated["accuracy"] < baseline_full["accuracy"] - 0.50
    )

    if weights.tobytes() != original_weight_bytes:
        raise AssertionError("search mutated the in-memory source weights")
    bit_difference = float32_winner.old_bits ^ float32_winner.new_bits
    if bit_difference != 1 << float32_winner.bit_index:
        raise AssertionError("selected float32 candidate differs by more than one bit")

    distribution_path = args.output_dir / "prediction-distributions.csv"
    per_class_path = args.output_dir / "per-class-accuracy.csv"
    per_example_path = args.output_dir / "winner-predictions.csv"
    write_distribution_csv(distribution_path, full_scenarios, len(labels))
    write_per_class_csv(per_class_path, full_scenarios)
    write_winner_predictions_csv(
        per_example_path,
        labels,
        model_logits,
        actual_mutated_logits,
        winner_feature_activations,
        float32_winner.target_index,
    )
    evidence_path = args.output_dir / "prediction-collapse-evidence.png"
    draw_evidence_chart(
        evidence_path,
        baseline_full,
        float32_mutated,
        float32_winner.target_class,
    )

    archive_path = args.data_dir / "cifar-10-python.tar.gz"
    test_batch_path = args.data_dir / "cifar-10-batches-py" / "test_batch"
    if not archive_path.exists() or not test_batch_path.exists():
        raise FileNotFoundError("expected CIFAR-10 archive and test batch after download")
    archive_md5 = md5_file(archive_path)
    if archive_md5 != CIFAR10_ARCHIVE_MD5:
        raise RuntimeError(f"CIFAR-10 archive MD5 mismatch: {archive_md5}")

    elapsed = time.time() - started
    results = {
        "schema_version": 1,
        "premise": "A single finite float32 bit flip in the final layer makes one animal class dominate at least 90% of predictions and drops top-1 accuracy by more than 50 percentage points.",
        "premise_passed": premise_passed,
        "selection_rule": "Among finite candidates targeting CIFAR-10 animal classes, maximize calibration target share, then minimize calibration accuracy, then choose the smaller flat weight index and bit index.",
        "scope": {
            "model": MODEL_NAME,
            "searched_parameter": "fc.weight",
            "shape": list(weights.shape),
            "calibration_images": int(len(calibration_indices)),
            "calibration_images_per_class": args.calibration_per_class,
            "full_test_images": int(len(labels)),
            "animal_classes": [CIFAR10_CLASSES[index] for index in ANIMAL_CLASS_INDICES],
            "float32_candidates": int(weights.size * 32),
            "float16_candidates": int(weights.size * 16),
            "int8_candidates": int(weights.size * 8),
            "total_candidates": int(weights.size * (32 + 16 + 8)),
            "precision_comparison_boundary": "Only final-layer weight storage changes; the ResNet backbone and captured features remain float32.",
        },
        "model": {
            "repository": MODEL_REPOSITORY,
            "source_commit": MODEL_SOURCE_COMMIT,
            "hub_spec": MODEL_HUB_SPEC,
            "checkpoint_url": f"{MODEL_REPOSITORY}/releases/download/resnet/{CHECKPOINT_FILENAME}",
            "checkpoint_sha256": sha256_file(checkpoint_path),
            "checkpoint_bytes": checkpoint_path.stat().st_size,
            "parameter_count": int(sum(parameter.numel() for parameter in model.parameters())),
        },
        "dataset": {
            "name": "CIFAR-10 test batch",
            "official_page": "https://www.cs.toronto.edu/~kriz/cifar.html",
            "download_mirror": CIFAR10_MIRROR_URL,
            "archive_md5": archive_md5,
            "archive_sha256": sha256_file(archive_path),
            "test_batch_sha256": sha256_file(test_batch_path),
            "classes": CIFAR10_CLASSES,
            "normalization_mean": NORMALIZE_MEAN,
            "normalization_std": NORMALIZE_STD,
        },
        "baseline_float32_full_test": json_ready_metrics(baseline_full),
        "calibration_format_baselines": format_baselines,
        "winners_full_test": full_winner_details,
        "field_summaries": summaries,
        "animal_target_summaries": animal_target_summaries,
        "controls": {
            "captured_features_reconstruct_model_logits": True,
            "original_model_and_reconstructed_predictions_match": bool(
                np.array_equal(model_logits.argmax(axis=1), reconstructed_logits.argmax(axis=1))
            ),
            "live_mutated_model_matches_analytical_predictions": True,
            "favorite_predictions_exactly_match_positive_winner_feature": True,
            "search_left_source_weights_unchanged": True,
            "live_model_parameter_restored_to_original_bits": True,
            "selected_float32_old_xor_new": f"0x{bit_difference:08x}",
            "selected_float32_exactly_one_bit": True,
        },
        "environment": {
            "python": sys.version.split()[0],
            "platform": platform.platform(),
            "torch": torch.__version__,
            "torchvision": torchvision.__version__,
            "numpy": np.__version__,
            "device": str(device),
            "gpu": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
            "elapsed_seconds": elapsed,
        },
    }
    results_path = args.output_dir / "results.json"
    results_path.write_bytes((json.dumps(results, indent=2) + "\n").encode("utf-8"))

    output_hashes = {}
    for path in sorted(args.output_dir.iterdir()):
        if path.name == "manifest.json" or not path.is_file():
            continue
        output_hashes[path.name] = {
            "sha256": sha256_file(path),
            "bytes": path.stat().st_size,
        }
    manifest = {
        "schema_version": 1,
        "premise_passed": premise_passed,
        "favorite_animal": float32_winner.target_class,
        "files": output_hashes,
    }
    (args.output_dir / "manifest.json").write_bytes(
        (json.dumps(manifest, indent=2) + "\n").encode("utf-8")
    )

    print(json.dumps({
        "premise_passed": premise_passed,
        "favorite_animal": float32_winner.target_class,
        "baseline_accuracy": baseline_full["accuracy"],
        "mutated_accuracy": float32_mutated["accuracy"],
        "mutated_favorite_share": float32_mutated["prediction_shares"][float32_winner.target_index],
        "outputs": str(args.output_dir),
        "elapsed_seconds": elapsed,
    }, indent=2))
    return 0 if premise_passed else 2


if __name__ == "__main__":
    raise SystemExit(main())
