"""Compute-matched comparison of PGN-style flat-maxima attack vs MI-FGSM.

Question: at a fixed gradient-call budget (not fixed iteration count), how much
of PGN's reported transferability gain over MI-FGSM survives?

Run: see run.sh. Use --quick for a tiny smoke test of the full pipeline.
"""
import argparse
import io
import json
import os
import sys
import time

os.environ.setdefault("OMP_NUM_THREADS", "2")

import numpy as np
import pandas as pd
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image

torch.set_num_threads(2)

EPS = 8.0 / 255.0
CIFAR_MEAN = (0.4914, 0.4822, 0.4465)
CIFAR_STD = (0.2471, 0.2435, 0.2616)

SURROGATE_NAME = "cifar10_resnet20"
TARGET_NAMES = [
    "cifar10_vgg11_bn",
    "cifar10_mobilenetv2_x1_0",
    "cifar10_shufflenetv2_x1_0",
    "cifar10_repvgg_a0",
]


class NormalizedModel(nn.Module):
    """Wraps a CIFAR model so it accepts raw [0,1] pixel input."""

    def __init__(self, model):
        super().__init__()
        self.model = model
        mean = torch.tensor(CIFAR_MEAN).view(1, 3, 1, 1)
        std = torch.tensor(CIFAR_STD).view(1, 3, 1, 1)
        self.register_buffer("mean", mean)
        self.register_buffer("std", std)

    def forward(self, x):
        return self.model((x - self.mean) / self.std)


def log(msg):
    print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True)


def load_model(name, retries=3):
    last_err = None
    for attempt in range(retries):
        try:
            m = torch.hub.load(
                "chenyaofo/pytorch-cifar-models", name, pretrained=True, trust_repo=True
            )
            m.eval()
            for p in m.parameters():
                p.requires_grad_(False)
            return NormalizedModel(m)
        except Exception as e:  # noqa: BLE001 - fallback path per task spec
            last_err = e
            log(f"torch.hub.load({name}) attempt {attempt+1} failed: {e}; retrying")
            time.sleep(2)
    raise RuntimeError(f"Failed to load {name} after {retries} attempts") from last_err


def load_cifar10_test(parquet_path):
    df = pd.read_parquet(parquet_path)
    imgs = np.stack(
        [np.array(Image.open(io.BytesIO(r["bytes"]))) for r in df["img"]]
    ).astype(np.float32) / 255.0
    labels = df["label"].values.astype(np.int64)
    x = torch.from_numpy(imgs).permute(0, 3, 1, 2).contiguous()  # N,C,H,W in [0,1]
    y = torch.from_numpy(labels)
    return x, y


def select_correctly_classified(x_pool, y_pool, models, n, seed, batch_size=256):
    """Pick n images correctly classified by every model in `models`, deterministic per seed."""
    rng = np.random.RandomState(seed)
    order = rng.permutation(len(x_pool))
    chosen_idx = []
    scan_count = 0
    with torch.no_grad():
        for start in range(0, len(order), batch_size):
            if len(chosen_idx) >= n:
                break
            idx = order[start:start + batch_size]
            xb = x_pool[idx]
            yb = y_pool[idx]
            ok = torch.ones(len(idx), dtype=torch.bool)
            for m in models:
                pred = m(xb).argmax(1)
                ok &= (pred == yb)
            scan_count += len(idx)
            for local_i, global_i in enumerate(idx):
                if ok[local_i]:
                    chosen_idx.append(int(global_i))
                    if len(chosen_idx) >= n:
                        break
    if len(chosen_idx) < n:
        raise RuntimeError(f"Only found {len(chosen_idx)}/{n} qualifying images after scanning {scan_count}")
    chosen_idx = chosen_idx[:n]
    log(f"  seed={seed}: scanned {scan_count} candidates to find {n} images correct on all {len(models)} models")
    return x_pool[chosen_idx].clone(), y_pool[chosen_idx].clone()


def l1_normalize(grad):
    n = grad.abs().sum(dim=[1, 2, 3], keepdim=True).clamp_min(1e-12)
    return grad / n


def mi_fgsm(model, x0, y, eps, T, decay=1.0):
    x_adv = x0.clone()
    g = torch.zeros_like(x0)
    alpha = 2.5 * eps / T
    grad_calls = 0
    for _ in range(T):
        x_adv.requires_grad_(True)
        out = model(x_adv)
        loss = F.cross_entropy(out, y)
        grad = torch.autograd.grad(loss, x_adv)[0]
        grad_calls += 1
        g = decay * g + l1_normalize(grad)
        x_adv = x_adv.detach() + alpha * g.sign()
        x_adv = torch.min(torch.max(x_adv, x0 - eps), x0 + eps)
        x_adv = torch.clamp(x_adv, 0.0, 1.0)
    return x_adv.detach(), grad_calls


def pgn_flat_maxima(model, x0, y, eps, T, N=3, zeta=3.0, balance=0.5, chi=None, decay=1.0):
    if chi is None:
        chi = eps
    x_adv = x0.clone()
    g = torch.zeros_like(x0)
    alpha = 2.5 * eps / T
    radius = zeta * eps
    grad_calls = 0
    for _ in range(T):
        combined_sum = torch.zeros_like(x0)
        for _ in range(N):
            noise = (torch.rand_like(x0) * 2 - 1) * radius
            x_s = (x_adv + noise).clamp(0.0, 1.0).detach().requires_grad_(True)
            out1 = model(x_s)
            loss1 = F.cross_entropy(out1, y)
            grad1 = torch.autograd.grad(loss1, x_s)[0]
            grad_calls += 1

            x_p = (x_s.detach() + chi * grad1.sign()).clamp(0.0, 1.0).detach().requires_grad_(True)
            out2 = model(x_p)
            loss2 = F.cross_entropy(out2, y)
            grad2 = torch.autograd.grad(loss2, x_p)[0]
            grad_calls += 1

            combined_sum = combined_sum + (1 - balance) * grad1 + balance * grad2
        avg_grad = combined_sum / N
        g = decay * g + l1_normalize(avg_grad)
        x_adv = x_adv.detach() + alpha * g.sign()
        x_adv = torch.min(torch.max(x_adv, x0 - eps), x0 + eps)
        x_adv = torch.clamp(x_adv, 0.0, 1.0)
    return x_adv.detach(), grad_calls


@torch.no_grad()
def success_rate(model, x_adv, y):
    pred = model(x_adv).argmax(1)
    return (pred != y).float().mean().item()


def neighborhood_grad_norm(model, x_adv, y, radius, n_samples):
    norms = []
    for _ in range(n_samples):
        noise = (torch.rand_like(x_adv) * 2 - 1) * radius
        x_s = (x_adv + noise).clamp(0.0, 1.0).detach().requires_grad_(True)
        out = model(x_s)
        loss = F.cross_entropy(out, y, reduction="sum")
        grad = torch.autograd.grad(loss, x_s)[0]
        l2 = grad.flatten(1).norm(dim=1)
        norms.append(l2)
    return torch.stack(norms, dim=0).mean().item()


def mean_std(values):
    arr = np.array(values, dtype=np.float64)
    return float(arr.mean()), float(arr.std())


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--grad-budgets", type=str, default="36,72")
    ap.add_argument("--seeds", type=str, default="0,1,2")
    ap.add_argument("--n-images", type=int, default=128)
    ap.add_argument("--quick", action="store_true", help="tiny smoke test of the pipeline")
    ap.add_argument("--out", type=str, default="results.json")
    ap.add_argument("--figs-dir", type=str, default="figs")
    ap.add_argument("--data", type=str, default="data/cifar10_test.parquet")
    args = ap.parse_args()

    t_start = time.time()

    if args.quick:
        grad_budgets = [12]
        seeds = [0]
        n_images = 16
    else:
        grad_budgets = [int(v) for v in args.grad_budgets.split(",")]
        seeds = [int(v) for v in args.seeds.split(",")]
        n_images = args.n_images

    log(f"Config: grad_budgets={grad_budgets} seeds={seeds} n_images={n_images} eps={EPS:.5f} quick={args.quick}")

    log("Loading models...")
    surrogate = load_model(SURROGATE_NAME)
    targets = {name: load_model(name) for name in TARGET_NAMES}
    all_models = [surrogate] + list(targets.values())
    log(f"Loaded surrogate={SURROGATE_NAME}, targets={TARGET_NAMES}")

    log(f"Loading CIFAR-10 test set from {args.data} ...")
    x_pool, y_pool = load_cifar10_test(args.data)
    log(f"Pool size: {len(x_pool)}")

    ARM_DEFS = [
        ("baseline-mi-equal-steps", True, "mi"),
        ("baseline-mi-equal-budget", True, "mi"),
        ("pgn-flat-maxima", False, "pgn"),
    ]

    # raw[(arm_base_name, grad_budget)] = {"seeds": [...], metric_name: [...per seed...]}
    raw = {}
    nb_radius = 2.0 / 255.0

    for gb in grad_budgets:
        for seed in seeds:
            log(f"=== grad_budget={gb} seed={seed} ===")
            torch.manual_seed(seed)
            np.random.seed(seed)
            x0, y0 = select_correctly_classified(x_pool, y_pool, all_models, n_images, seed)
            log(f"  selected batch: x0={tuple(x0.shape)} y0={tuple(y0.shape)}")

            for arm_name, is_baseline, kind in ARM_DEFS:
                t0 = time.time()
                if arm_name == "baseline-mi-equal-steps":
                    T = max(1, gb // 6)
                    x_adv, grad_calls = mi_fgsm(surrogate, x0, y0, EPS, T)
                elif arm_name == "baseline-mi-equal-budget":
                    T = gb
                    x_adv, grad_calls = mi_fgsm(surrogate, x0, y0, EPS, T)
                elif arm_name == "pgn-flat-maxima":
                    T = max(1, gb // 6)
                    x_adv, grad_calls = pgn_flat_maxima(surrogate, x0, y0, EPS, T, N=3, zeta=3.0, balance=0.5, chi=EPS)
                else:
                    raise ValueError(arm_name)

                wb = success_rate(surrogate, x_adv, y0)
                per_target = {name: success_rate(m, x_adv, y0) for name, m in targets.items()}
                transfer_mean = float(np.mean(list(per_target.values())))
                ngn = neighborhood_grad_norm(surrogate, x_adv, y0, nb_radius, 5)

                dt = time.time() - t0
                log(
                    f"  arm={arm_name:28s} T={T:3d} grad_calls={grad_calls:3d} "
                    f"whitebox={wb:.3f} transfer_mean={transfer_mean:.3f} "
                    f"ngn={ngn:.4f} per_target={ {k: round(v,3) for k,v in per_target.items()} } "
                    f"time={dt:.1f}s"
                )

                key = (arm_name, gb)
                if key not in raw:
                    raw[key] = {
                        "seeds": [], "grad_calls": [], "whitebox_success_rate": [],
                        "transfer_success_rate": [], "neighborhood_grad_norm": [],
                        "per_target": {name: [] for name in TARGET_NAMES},
                    }
                raw[key]["seeds"].append(seed)
                raw[key]["grad_calls"].append(grad_calls)
                raw[key]["whitebox_success_rate"].append(wb)
                raw[key]["transfer_success_rate"].append(transfer_mean)
                raw[key]["neighborhood_grad_norm"].append(ngn)
                for name in TARGET_NAMES:
                    raw[key]["per_target"][name].append(per_target[name])

    # ---- assemble results.json ----
    arms_out = []
    arm_display = {
        "baseline-mi-equal-steps": (
            "MI-FGSM (decay=1.0), T=grad_budget/6 -- same iteration count as PGN but only "
            "grad_budget/6 gradient calls (the unfair, literature-standard comparison)."
        ),
        "baseline-mi-equal-budget": (
            "MI-FGSM (decay=1.0), T=grad_budget -- consumes exactly the same number of "
            "fwd+bwd calls as PGN (the compute-matched, fair comparison)."
        ),
        "pgn-flat-maxima": (
            "PGN-style flat-maxima attack: per iteration, 3 neighborhood samples each "
            "contribute 2 gradient calls (sample point + sign-step prediction point), "
            "weighted 0.5/0.5 and averaged, momentum decay=1.0. 6 grad calls/iteration, "
            "T=grad_budget/6 iterations -> grad_budget total calls."
        ),
    }

    for arm_name, gb in sorted(raw.keys(), key=lambda k: (k[0], k[1])):
        d = raw[(arm_name, gb)]
        metrics = {}
        for metric_name in ["whitebox_success_rate", "transfer_success_rate", "neighborhood_grad_norm"]:
            mean, std = mean_std(d[metric_name])
            metrics[metric_name] = {"per_seed": d[metric_name], "mean": mean, "std": std}
        gc_mean, gc_std = mean_std(d["grad_calls"])
        metrics["fwd_bwd_calls"] = {"per_seed": d["grad_calls"], "mean": gc_mean, "std": gc_std}
        for name in TARGET_NAMES:
            mean, std = mean_std(d["per_target"][name])
            metrics[f"transfer_success_rate__{name}"] = {
                "per_seed": d["per_target"][name], "mean": mean, "std": std,
            }
        arms_out.append({
            "name": f"{arm_name}-gb{gb}",
            "is_baseline": arm_name.startswith("baseline"),
            "what": f"{arm_display[arm_name]} grad_budget={gb}.",
            "metrics": metrics,
        })

    # ---- falsification logic (pre-registered rule, applied mechanically) ----
    def arm_metric_per_run(arm_name, metric):
        vals = []
        for gb in grad_budgets:
            vals.extend(raw[(arm_name, gb)][metric])
        return vals

    tsr_pgn = arm_metric_per_run("pgn-flat-maxima", "transfer_success_rate")
    tsr_eqstep = arm_metric_per_run("baseline-mi-equal-steps", "transfer_success_rate")
    tsr_eqbudget = arm_metric_per_run("baseline-mi-equal-budget", "transfer_success_rate")

    delta_eqstep = [100 * (p - b) for p, b in zip(tsr_pgn, tsr_eqstep)]
    delta_eqbudget = [100 * (p - b) for p, b in zip(tsr_pgn, tsr_eqbudget)]
    mean_delta_eqstep = float(np.mean(delta_eqstep))
    mean_delta_eqbudget = float(np.mean(delta_eqbudget))
    all_positive_eqbudget = all(d > 0 for d in delta_eqbudget)

    part_a_holds = mean_delta_eqstep >= 5.0
    if not part_a_holds:
        outcome = "inconclusive"
        outcome_reason = (
            f"Part A premise failed: equal-steps advantage mean={mean_delta_eqstep:.2f}pp < 5pp, "
            "so the literature-bias premise the prediction is built on did not even replicate here."
        )
    elif mean_delta_eqbudget <= 2.0:
        outcome = "confirmed"
        outcome_reason = (
            f"Equal-steps advantage={mean_delta_eqstep:.2f}pp (>=5pp, part A holds); "
            f"equal-budget advantage shrank to {mean_delta_eqbudget:.2f}pp (<=2pp threshold)."
        )
    elif mean_delta_eqbudget >= 5.0 and all_positive_eqbudget:
        outcome = "refuted"
        outcome_reason = (
            f"Equal-steps advantage={mean_delta_eqstep:.2f}pp (>=5pp, part A holds); "
            f"equal-budget advantage stayed at {mean_delta_eqbudget:.2f}pp (>=5pp) and every "
            f"one of {len(delta_eqbudget)} (budget,seed) runs was positive: {delta_eqbudget}."
        )
    else:
        outcome = "inconclusive"
        outcome_reason = (
            f"Equal-steps advantage={mean_delta_eqstep:.2f}pp (part A holds); equal-budget "
            f"advantage={mean_delta_eqbudget:.2f}pp falls between the 2pp confirm and 5pp/all-positive "
            f"refute thresholds (per-run deltas: {delta_eqbudget})."
        )

    log(f"Falsification check: {outcome_reason}")

    headline_gb = grad_budgets[-1]  # larger budget = more stable estimate
    headline_baseline_arm = next(a for a in arms_out if a["name"] == f"baseline-mi-equal-budget-gb{headline_gb}")
    headline_proposed_arm = next(a for a in arms_out if a["name"] == f"pgn-flat-maxima-gb{headline_gb}")
    baseline_mean = headline_baseline_arm["metrics"]["transfer_success_rate"]["mean"]
    proposed_mean = headline_proposed_arm["metrics"]["transfer_success_rate"]["mean"]

    claim_map = {
        "confirmed": (
            f"At grad_budget={headline_gb} (compute-matched), MI-FGSM ({baseline_mean*100:.1f}%) and "
            f"PGN-flat-maxima ({proposed_mean*100:.1f}%) transfer success rates are close "
            f"(pooled mean advantage {mean_delta_eqbudget:.1f}pp across both budgets/seeds, vs "
            f"{mean_delta_eqstep:.1f}pp under the unfair equal-iteration comparison): most of PGN's "
            f"claimed transferability edge over MI-FGSM is a compute-budget artifact, not a flat-maxima effect."
        ),
        "refuted": (
            f"At grad_budget={headline_gb} (compute-matched), PGN-flat-maxima ({proposed_mean*100:.1f}%) "
            f"still beats MI-FGSM ({baseline_mean*100:.1f}%) by {mean_delta_eqbudget:.1f}pp pooled across "
            f"both budgets/seeds, with every individual run positive: the transferability gain survives "
            f"compute-matching and is not just an artifact of extra gradient calls."
        ),
        "inconclusive": (
            f"At grad_budget={headline_gb}, MI-FGSM={baseline_mean*100:.1f}%, PGN-flat-maxima="
            f"{proposed_mean*100:.1f}% transfer success; {outcome_reason}"
        ),
    }

    runtime_sec = time.time() - t_start

    results = {
        "task_slug": "flat-maxima-transferability-compute-matched",
        "question": (
            "在固定梯度调用预算(而非固定迭代数)下,PGN 式平坦极值攻击相对 MI-FGSM 的迁移性增益还剩多少?"
            "平坦度本身是否是中介变量?"
        ),
        "falsifiable_prediction": (
            "预测:相同迭代数下平坦化臂对 MI-FGSM 的平均迁移成功率优势 >=5 个百分点;改为相同梯度调用预算后"
            "优势缩到 <=2 个百分点或反转。若在相同梯度调用预算下平坦化臂在两个预算档 x 3 个种子的平均值仍"
            "领先 >=5 个百分点且每个种子都为正,则预测被推翻。"
        ),
        "prediction_outcome": outcome,
        "negative_result": outcome == "refuted",
        "dataset": (
            f"CIFAR-10 test set ({len(x_pool)} images, torchvision-format via HF parquet mirror); "
            f"per (grad_budget, seed) combo: {n_images} images sampled and filtered to be correctly "
            f"classified by surrogate cifar10_resnet20 AND all 4 targets "
            f"(cifar10_vgg11_bn, cifar10_mobilenetv2_x1_0, cifar10_shufflenetv2_x1_0, cifar10_repvgg_a0)."
        ),
        "env": {
            "python": sys.version.split()[0],
            "torch": torch.__version__,
            "extra_packages": ["pyarrow==25.0.0"],
        },
        "seeds": seeds,
        "arms": arms_out,
        "headline": {
            "metric": "transfer_success_rate",
            "baseline_mean": baseline_mean,
            "proposed_mean": proposed_mean,
            "delta": proposed_mean - baseline_mean,
            "claim": claim_map[outcome],
        },
        "deviations": [] if not args.quick else ["--quick smoke-test mode: n_images=16, 1 seed, grad_budget=12"],
        "runtime_sec": runtime_sec,
        "diagnostics": {
            "eps": EPS,
            "grad_budgets_swept": grad_budgets,
            "mean_delta_pp_equal_steps": mean_delta_eqstep,
            "mean_delta_pp_equal_budget": mean_delta_eqbudget,
            "per_run_delta_pp_equal_budget": delta_eqbudget,
            "all_positive_equal_budget": all_positive_eqbudget,
            "outcome_reason": outcome_reason,
            "pgn_hyperparams": {"N": 3, "zeta": 3.0, "balance": 0.5, "chi": EPS, "decay": 1.0},
        },
    }

    with open(args.out, "w") as f:
        json.dump(results, f, indent=2)
    log(f"Wrote {args.out}")

    # ---- figures ----
    os.makedirs(args.figs_dir, exist_ok=True)
    import matplotlib
    matplotlib.use("Agg")
    import matplotlib.pyplot as plt

    fig, axes = plt.subplots(1, 2, figsize=(11, 4.5))

    ax = axes[0]
    width = 0.25
    x_pos = np.arange(len(grad_budgets))
    arm_order = ["baseline-mi-equal-steps", "baseline-mi-equal-budget", "pgn-flat-maxima"]
    colors = {"baseline-mi-equal-steps": "#999999", "baseline-mi-equal-budget": "#4c72b0", "pgn-flat-maxima": "#c44e52"}
    for i, arm_name in enumerate(arm_order):
        means = [raw[(arm_name, gb)]["transfer_success_rate"] for gb in grad_budgets]
        m = [float(np.mean(v)) * 100 for v in means]
        s = [float(np.std(v)) * 100 for v in means]
        ax.bar(x_pos + (i - 1) * width, m, width, yerr=s, label=arm_name, color=colors[arm_name], capsize=3)
    ax.set_xticks(x_pos)
    ax.set_xticklabels([f"grad_budget={gb}" for gb in grad_budgets])
    ax.set_ylabel("Transfer success rate (%)")
    ax.set_title("Transfer success: equal-steps vs equal-budget vs PGN")
    ax.legend(fontsize=8)

    ax2 = axes[1]
    for arm_name in arm_order:
        xs, ys = [], []
        for gb in grad_budgets:
            xs.append(float(np.mean(raw[(arm_name, gb)]["neighborhood_grad_norm"])))
            ys.append(float(np.mean(raw[(arm_name, gb)]["transfer_success_rate"])) * 100)
        ax2.scatter(xs, ys, label=arm_name, color=colors[arm_name], s=80)
        for gb, x, y in zip(grad_budgets, xs, ys):
            ax2.annotate(f"gb{gb}", (x, y), fontsize=7, xytext=(3, 3), textcoords="offset points")
    ax2.set_xlabel("Neighborhood grad L2 norm (surrogate, lower=flatter)")
    ax2.set_ylabel("Transfer success rate (%)")
    ax2.set_title("Flatness vs transferability")
    ax2.legend(fontsize=8)

    fig.tight_layout()
    fig_path = os.path.join(args.figs_dir, "transfer_and_flatness.png")
    fig.savefig(fig_path, dpi=140)
    log(f"Wrote {fig_path}")

    log(f"Done in {runtime_sec:.1f}s. outcome={outcome}")


if __name__ == "__main__":
    main()
