{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":10338,"databundleVersionId":862042,"sourceType":"competition"},{"sourceId":23812,"sourceType":"datasetVersion","datasetId":17810},{"sourceId":14523676,"sourceType":"datasetVersion","datasetId":9275978},{"sourceId":14529603,"sourceType":"datasetVersion","datasetId":9280032}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =========================\n# PneumoniaKD: Calibration + Spec@Sens>=0.95 (RAW + TS), 5 seeds\n# FIXED: medmnist transform robustness + checkpoint-matched TinyCNN (classifier head + BN variant)\n# =========================\n\nimport os, math\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\nfrom sklearn.metrics import roc_curve\n\n# ---- Install/import MedMNIST ----\n!pip -q install medmnist\n\nimport medmnist\nfrom medmnist import PneumoniaMNIST\nfrom torchvision import transforms\nfrom torchvision.models import resnet18\nfrom PIL import Image\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ntorch.backends.cudnn.benchmark = True\nprint(\"DEVICE:\", DEVICE, \"torch:\", torch.__version__, \"medmnist:\", getattr(medmnist, \"__version__\", \"unknown\"))\n\n# ---- Set checkpoint directory ----\nCKPT_DIR = \"/kaggle/input/okok1ok/publish_pack_v2\"  # <-- CHANGE THIS\n\nMODELS = {\n    \"teacher\":  {\"pattern\": \"Teacher_seed{seed}.pt\"},\n    \"vanilla\":  {\"pattern\": \"Vanilla_seed{seed}.pt\"},\n    \"kd_resp\":  {\"pattern\": \"KDResp_seed{seed}.pt\"},\n    \"kd_fitnets\": {\"pattern\": \"KDFit_seed{seed}.pt\"},\n}\nSEEDS = [0,1,2,3,4]\n\n# ---- Robust PIL conversion (medmnist can yield PIL or numpy) ----\ndef to_pil(x):\n    if isinstance(x, Image.Image):\n        return x\n    if isinstance(x, np.ndarray):\n        if x.ndim == 3 and x.shape[-1] == 1:\n            x = x[..., 0]\n        return Image.fromarray(x)\n    if torch.is_tensor(x):\n        arr = x.detach().cpu().numpy()\n        if arr.ndim == 3 and arr.shape[0] in (1,3):\n            arr = np.transpose(arr, (1,2,0))\n        if arr.dtype != np.uint8:\n            arr = (255*(arr - arr.min())/(arr.max()-arr.min()+1e-8)).clip(0,255).astype(np.uint8)\n        if arr.ndim == 3 and arr.shape[-1] == 1:\n            arr = arr[...,0]\n        return Image.fromarray(arr)\n    return Image.open(x)\n\n# ---- Data ----\nimagenet_mean = (0.485, 0.456, 0.406)\nimagenet_std  = (0.229, 0.224, 0.225)\n\ntransform = transforms.Compose([\n    transforms.Lambda(to_pil),\n    transforms.Lambda(lambda img: img.convert(\"RGB\")),\n    transforms.Resize((224,224)),\n    transforms.ToTensor(),\n    transforms.Normalize(imagenet_mean, imagenet_std),\n])\n\nval_ds  = PneumoniaMNIST(split=\"val\",  transform=transform, download=True)\ntest_ds = PneumoniaMNIST(split=\"test\", transform=transform, download=True)\n\nVAL_LOADER  = DataLoader(val_ds,  batch_size=256, shuffle=False, num_workers=2, pin_memory=True)\nTEST_LOADER = DataLoader(test_ds, batch_size=256, shuffle=False, num_workers=2, pin_memory=True)\n\n# ---- Student model variants: use `classifier` head to match your checkpoints ----\nclass TinyCNN_vA(nn.Module):\n    # no BatchNorm, classifier head\n    def __init__(self, num_classes=2):\n        super().__init__()\n        self.features = nn.Sequential(\n            nn.Conv2d(3,16,3,padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2),\n            nn.Conv2d(16,32,3,padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2),\n            nn.Conv2d(32,64,3,padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2),\n            nn.Conv2d(64,64,3,padding=1), nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool2d(1),\n        )\n        self.classifier = nn.Linear(64, num_classes)\n\n    def forward(self, x):\n        x = self.features(x).flatten(1)\n        return self.classifier(x)\n\nclass TinyCNN_vB(nn.Module):\n    # BatchNorm at tail (matches features.9.weight/bias pattern you saw)\n    def __init__(self, num_classes=2):\n        super().__init__()\n        self.features = nn.Sequential(\n            nn.Conv2d(3,16,3,padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2),\n            nn.Conv2d(16,32,3,padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2),\n            nn.Conv2d(32,64,3,padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2),\n            nn.Conv2d(64,64,3,padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool2d(1),\n        )\n        self.classifier = nn.Linear(64, num_classes)\n\n    def forward(self, x):\n        x = self.features(x).flatten(1)\n        return self.classifier(x)\n\ndef build_teacher():\n    m = resnet18(weights=None)\n    m.fc = nn.Linear(m.fc.in_features, 2)\n    return m\n\ndef normalize_state_dict(sd):\n    # support common formats\n    if isinstance(sd, dict) and \"state_dict\" in sd:\n        sd = sd[\"state_dict\"]\n    sd = {k.replace(\"module.\", \"\"): v for k, v in sd.items()}\n    return sd\n\ndef pick_student_arch(sd):\n    ck_keys = set(sd.keys())\n    cands = [(\"vB\", TinyCNN_vB()), (\"vA\", TinyCNN_vA())]\n    for name, m in cands:\n        if ck_keys == set(m.state_dict().keys()):\n            return name, m\n    # helpful debug if it still doesn't match\n    diffs=[]\n    for name, m in cands:\n        mk=set(m.state_dict().keys())\n        diffs.append((name, len(ck_keys.symmetric_difference(mk))))\n    diffs=sorted(diffs, key=lambda x:x[1])\n    raise RuntimeError(f\"No TinyCNN variant matches ckpt keys. Closest={diffs[0]}. Example ckpt keys: {sorted(list(ck_keys))[:12]}\")\n\ndef load_ckpt_strict(model, path):\n    sd = torch.load(path, map_location=\"cpu\")\n    sd = normalize_state_dict(sd)\n    model.load_state_dict(sd, strict=True)\n    return model\n\n# ---- Metrics ----\ndef softmax_probs_np(logits_np):\n    # logits_np shape (N,2)\n    x = torch.tensor(logits_np, dtype=torch.float32)\n    return F.softmax(x, dim=1).numpy()\n\ndef nll_binary(probs_pos, y, eps=1e-7):\n    probs_pos = np.clip(probs_pos, eps, 1-eps)\n    y = y.astype(np.float32)\n    return float(np.mean(-(y*np.log(probs_pos) + (1-y)*np.log(1-probs_pos))))\n\ndef brier_binary(probs_pos, y):\n    y = y.astype(np.float32)\n    return float(np.mean((probs_pos - y)**2))\n\ndef ece_confidence(probs, y, n_bins=15):\n    conf = np.max(probs, axis=1)\n    pred = np.argmax(probs, axis=1)\n    acc = (pred == y).astype(np.float32)\n\n    bins = np.linspace(0.0, 1.0, n_bins+1)\n    ece = 0.0\n    N = len(y)\n    for i in range(n_bins):\n        lo, hi = bins[i], bins[i+1]\n        if i == 0:\n            mask = (conf >= lo) & (conf <= hi)\n        else:\n            mask = (conf > lo) & (conf <= hi)\n        if not np.any(mask):\n            continue\n        ece += (np.sum(mask)/N) * abs(np.mean(acc[mask]) - np.mean(conf[mask]))\n    return float(ece)\n\ndef spec_at_sens95(probs_pos, y):\n    fpr, tpr, thr = roc_curve(y, probs_pos)\n    mask = (tpr >= 0.95)\n    if not np.any(mask):\n        return float(\"nan\")\n    spec = 1.0 - fpr[mask]\n    return float(np.max(spec))\n\n# ---- Temperature scaling ----\nclass TemperatureScaler(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.log_T = nn.Parameter(torch.zeros(1))\n\n    def forward(self, logits):\n        T = torch.exp(self.log_T).clamp(min=1e-3, max=100.0)\n        return logits / T\n\ndef fit_temperature(val_logits_np, val_y_np, max_iter=200):\n    scaler = TemperatureScaler().to(DEVICE)\n    val_logits = torch.tensor(val_logits_np, dtype=torch.float32, device=DEVICE)\n    val_y = torch.tensor(val_y_np, dtype=torch.long, device=DEVICE)\n    crit = nn.CrossEntropyLoss()\n    opt = torch.optim.LBFGS(scaler.parameters(), lr=0.1, max_iter=max_iter)\n\n    def closure():\n        opt.zero_grad()\n        loss = crit(scaler(val_logits), val_y)\n        loss.backward()\n        return loss\n\n    opt.step(closure)\n    return float(torch.exp(scaler.log_T).detach().cpu().item())\n\n@torch.no_grad()\ndef collect_logits(model, loader):\n    model.eval()\n    all_logits, all_y = [], []\n    for x, y in loader:\n        x = x.to(DEVICE, non_blocking=True)\n        logits = model(x).detach().cpu().numpy()\n        y_np = y.squeeze().cpu().numpy().astype(np.int64)\n        all_logits.append(logits)\n        all_y.append(y_np)\n    return np.concatenate(all_logits, 0), np.concatenate(all_y, 0)\n\n# ---- Run ----\nrows = []\n\nfor seed in SEEDS:\n    # pick student architecture once per seed using vanilla checkpoint\n    vanilla_path = os.path.join(CKPT_DIR, MODELS[\"vanilla\"][\"pattern\"].format(seed=seed))\n    if not os.path.exists(vanilla_path):\n        raise FileNotFoundError(vanilla_path)\n    sd_any = normalize_state_dict(torch.load(vanilla_path, map_location=\"cpu\"))\n    arch_name, student_template = pick_student_arch(sd_any)\n    print(f\"[seed {seed}] student arch = {arch_name}\")\n\n    for tag, cfg in MODELS.items():\n        ckpt_path = os.path.join(CKPT_DIR, cfg[\"pattern\"].format(seed=seed))\n        if not os.path.exists(ckpt_path):\n            raise FileNotFoundError(f\"Missing checkpoint: {ckpt_path}\")\n\n        if tag == \"teacher\":\n            model = build_teacher()\n        else:\n            model = type(student_template)()  # same class\n        model = model.to(DEVICE)\n\n        model = load_ckpt_strict(model, ckpt_path).to(DEVICE)\n\n        val_logits, val_y = collect_logits(model, VAL_LOADER)\n        test_logits, test_y = collect_logits(model, TEST_LOADER)\n\n        # RAW\n        probs_raw = softmax_probs_np(test_logits)\n        p1_raw = probs_raw[:,1]\n        ece_raw = ece_confidence(probs_raw, test_y, n_bins=15)\n        nll_raw = nll_binary(p1_raw, test_y, eps=1e-7)\n        brier_raw = brier_binary(p1_raw, test_y)\n        spec95_raw = spec_at_sens95(p1_raw, test_y)\n\n        # TS\n        Tstar = fit_temperature(val_logits, val_y, max_iter=200)\n        probs_ts = softmax_probs_np(test_logits / Tstar)\n        p1_ts = probs_ts[:,1]\n        ece_ts = ece_confidence(probs_ts, test_y, n_bins=15)\n        nll_ts = nll_binary(p1_ts, test_y, eps=1e-7)\n        brier_ts = brier_binary(p1_ts, test_y)\n        spec95_ts = spec_at_sens95(p1_ts, test_y)\n\n        rows.append({\n            \"tag\": tag, \"seed\": seed,\n            \"student_arch\": (\"resnet18\" if tag==\"teacher\" else arch_name),\n            \"ckpt\": os.path.basename(ckpt_path),\n            \"T\": Tstar,\n            \"raw_ece_conf\": ece_raw, \"raw_nll\": nll_raw, \"raw_brier\": brier_raw, \"raw_spec_at_sens95\": spec95_raw,\n            \"ts_ece_conf\": ece_ts, \"ts_nll\": nll_ts, \"ts_brier\": brier_ts, \"ts_spec_at_sens95\": spec95_ts,\n        })\n\n        print(\"done\", tag, \"seed\", seed, \"T*\", round(Tstar,3))\n\ndf = pd.DataFrame(rows)\ndf.to_csv(\"calibration_per_seed_WITH_SPEC.csv\", index=False)\nprint(\"Wrote calibration_per_seed_WITH_SPEC.csv\")\n\n# ---- Summary ----\ndef summarize(group, col):\n    x = group[col].astype(float).values\n    n = len(x)\n    m = float(np.mean(x))\n    sd = float(np.std(x, ddof=1)) if n > 1 else 0.0\n    try:\n        import scipy.stats as st\n        tcrit = st.t.ppf(0.975, df=n-1) if n > 1 else 0.0\n    except:\n        tcrit = 1.96 if n > 1 else 0.0\n    half = tcrit * sd / math.sqrt(n) if n > 1 else 0.0\n    return pd.Series({\"mean\": m, \"sd\": sd, \"ci_lo\": m-half, \"ci_hi\": m+half, \"n\": n})\n\nsummary_cols = [\"T\",\"raw_ece_conf\",\"raw_nll\",\"raw_brier\",\"raw_spec_at_sens95\",\n                \"ts_ece_conf\",\"ts_nll\",\"ts_brier\",\"ts_spec_at_sens95\"]\n\nsumm = []\nfor tag, g in df.groupby(\"tag\"):\n    for col in summary_cols:\n        s = summarize(g, col)\n        s[\"tag\"] = tag\n        s[\"metric\"] = col\n        summ.append(s)\nsumm_df = pd.DataFrame(summ)[[\"tag\",\"metric\",\"mean\",\"sd\",\"ci_lo\",\"ci_hi\",\"n\"]]\nsumm_df.to_csv(\"calibration_summary_ci_WITH_SPEC.csv\", index=False)\nprint(\"Wrote calibration_summary_ci_WITH_SPEC.csv\")\n\n# --- FIXED LaTeX emission (use [\"mean\"] not .mean) ---\n\ndef fmt_ci(m, lo, hi, nd=4):\n    m = float(m); lo = float(lo); hi = float(hi)\n    return f\"{m:.{nd}f} [{lo:.{nd}f}, {hi:.{nd}f}]\"\n\ndef get_row(tag, metric):\n    r = summ_df[(summ_df.tag==tag) & (summ_df.metric==metric)].iloc[0]\n    return r\n\norder = [(\"teacher\",\"ResNet-18 (teacher)\"),\n         (\"vanilla\",\"TinyCNN (vanilla)\"),\n         (\"kd_resp\",\"TinyCNN + KD (response)\"),\n         (\"kd_fitnets\",\"TinyCNN + KD (FitNets)\")]\n\nraw_lines=[]\nts_lines=[]\nfor tag, name in order:\n    e = get_row(tag,\"raw_ece_conf\")\n    n = get_row(tag,\"raw_nll\")\n    b = get_row(tag,\"raw_brier\")\n    s = get_row(tag,\"raw_spec_at_sens95\")\n    raw_lines.append(\n        f\"{name} & {fmt_ci(e['mean'],e['ci_lo'],e['ci_hi'])} & {fmt_ci(n['mean'],n['ci_lo'],n['ci_hi'])} & {fmt_ci(b['mean'],b['ci_lo'],b['ci_hi'])} & {fmt_ci(s['mean'],s['ci_lo'],s['ci_hi'])} \\\\\\\\\"\n    )\n\n    Tm = get_row(tag,\"T\")\n    e2 = get_row(tag,\"ts_ece_conf\")\n    n2 = get_row(tag,\"ts_nll\")\n    b2 = get_row(tag,\"ts_brier\")\n    s2 = get_row(tag,\"ts_spec_at_sens95\")\n    ts_lines.append(\n        f\"{name} & {fmt_ci(Tm['mean'],Tm['ci_lo'],Tm['ci_hi'],nd=2)} & {fmt_ci(e2['mean'],e2['ci_lo'],e2['ci_hi'])} & {fmt_ci(n2['mean'],n2['ci_lo'],n2['ci_hi'])} & {fmt_ci(b2['mean'],b2['ci_lo'],b2['ci_hi'])} & {fmt_ci(s2['mean'],s2['ci_lo'],s2['ci_hi'])} \\\\\\\\\"\n    )\n\nlatex = f\"\"\"\n% ===== Calibration tables (generated) =====\n% ECE: 15 equal-width bins; confidence of predicted class vs correctness\n% NLL: probability clipping to [1e-7, 1-1e-7]\n\\\\begin{{table*}}[t]\n\\\\centering\n\\\\caption{{Calibration before temperature scaling (raw probabilities) on the test set (mean; 95\\\\% CI of mean, 5 seeds). Lower is better.}}\n\\\\label{{tab:cal_raw}}\n\\\\resizebox{{\\\\textwidth}}{{!}}{{%\n\\\\begin{{tabular}}{{lcccc}}\n\\\\toprule\n\\\\textbf{{Model}} & \\\\textbf{{ECE$_{{raw}}$}} & \\\\textbf{{NLL$_{{raw}}$}} & \\\\textbf{{Brier$_{{raw}}$}} & \\\\textbf{{Spec@Sens$\\\\ge$0.95 (raw)}} \\\\\\\\\n\\\\midrule\n{chr(10).join(raw_lines)}\n\\\\bottomrule\n\\\\end{{tabular}}}}\n\\\\end{{table*}}\n\n\\\\begin{{table*}}[t]\n\\\\centering\n\\\\caption{{Calibration after temperature scaling (TS) on the test set (mean; 95\\\\% CI of mean, 5 seeds). $T^*$ is fit on the validation split. Lower is better.}}\n\\\\label{{tab:cal}}\n\\\\resizebox{{\\\\textwidth}}{{!}}{{%\n\\\\begin{{tabular}}{{lccccc}}\n\\\\toprule\n\\\\textbf{{Model}} & \\\\textbf{{$T^*$}} & \\\\textbf{{ECE$_{{TS}}$}} & \\\\textbf{{NLL$_{{TS}}$}} & \\\\textbf{{Brier$_{{TS}}$}} & \\\\textbf{{Spec@Sens$\\\\ge$0.95 (TS)}} \\\\\\\\\n\\\\midrule\n{chr(10).join(ts_lines)}\n\\\\bottomrule\n\\\\end{{tabular}}}}\n\\\\end{{table*}}\n\"\"\"\n\nwith open(\"calibration_tables_WITH_SPEC.tex\",\"w\") as f:\n    f.write(latex)\n\nprint(\"Wrote calibration_tables_WITH_SPEC.tex\")\n\n\nprint(\"\\nDone.\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-17T09:29:35.698752Z","iopub.execute_input":"2026-01-17T09:29:35.699364Z","execution_failed":"2026-01-17T09:29:36.403Z"}},"outputs":[],"execution_count":null}]}