{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader, Sampler\nfrom torchvision import models, transforms\nfrom PIL import Image\nfrom sklearn.metrics import roc_auc_score\nfrom sklearn.model_selection import train_test_split\nfrom tqdm.auto import tqdm\n\n# ------------------------------------------------------------------\n# CONFIG\n# ------------------------------------------------------------------\nINPUT_DIRS = [\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part1\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part2\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-negative-part3\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs-positive-part-1\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs-part3\",\n    \"/kaggle/input/datasets/anushakirand/rsna-ich-preprocessed-pngs\",\n]\n\nTRAIN_CSV = \"/kaggle/input/competitions/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection/stage_2_train.csv\"\n\nMIN_FILE_SIZE_KB = 40\nMIN_FILE_SIZE_BYTES = MIN_FILE_SIZE_KB * 1024\n\nLABEL_COLS = [\"any\", \"epidural\", \"intraparenchymal\",\n              \"intraventricular\", \"subarachnoid\", \"subdural\"]\n\nBATCH_SIZE = 32\nNUM_EPOCHS = 15          # we'll test with 1 first before committing to full run\nLR = 1e-4\n\nVAL_SPLIT = 0.15\nTEST_SPLIT = 0.15\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nUSE_AMP = DEVICE.type == \"cuda\"\n\nEARLY_STOPPING_PATIENCE = 5\nCHECKPOINT_DIR = \"/kaggle/working/checkpoints_vit\"   # separate from teammate's ResNet dir\nTEST_IDS_PATH = \"/kaggle/working/vit_test_ids.csv\"\n\nNUM_WORKERS = 4\n\nprint(\"device:\", DEVICE, \"| AMP:\", USE_AMP)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:39:51.765698Z","iopub.execute_input":"2026-07-26T20:39:51.766486Z","iopub.status.idle":"2026-07-26T20:40:01.083748Z","shell.execute_reply.started":"2026-07-26T20:39:51.766453Z","shell.execute_reply":"2026-07-26T20:40:01.082935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def scan_input_dirs(input_dirs, min_size_bytes):\n    id_to_path = {}\n    skipped_small = 0\n    for d in input_dirs:\n        if not os.path.isdir(d):\n            print(f\"WARNING: directory not found, skipping: {d}\")\n            continue\n        filenames = [f for f in os.listdir(d) if f.endswith(\".png\")]\n        for f in tqdm(filenames, desc=f\"scanning {os.path.basename(d)}\"):\n            full_path = os.path.join(d, f)\n            if os.path.getsize(full_path) <= min_size_bytes:\n                skipped_small += 1\n                continue\n            img_id = f[:-4]   # strip \".png\"\n            id_to_path[img_id] = full_path\n    print(f\"Scanned {len(input_dirs)} directories: {len(id_to_path)} usable PNGs \"\n          f\"found, {skipped_small} filtered out as low-information (<= \"\n          f\"{MIN_FILE_SIZE_KB}KB).\")\n    return id_to_path\n\n\ndef load_labels(csv_path):\n    y = pd.read_csv(csv_path)\n    id_split = y.ID.str.rsplit(\"_\", n=1, expand=True)\n    y = pd.concat([id_split, y.Label], axis=1)\n    y.columns = [\"id\", \"sub_type\", \"label\"]\n    y = y.drop_duplicates(subset=[\"id\", \"sub_type\"])\n    df = y.pivot(index=\"id\", columns=\"sub_type\", values=\"label\")\n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:40:01.085185Z","iopub.execute_input":"2026-07-26T20:40:01.085562Z","iopub.status.idle":"2026-07-26T20:40:01.09394Z","shell.execute_reply.started":"2026-07-26T20:40:01.085538Z","shell.execute_reply":"2026-07-26T20:40:01.09322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"id_to_path = scan_input_dirs(INPUT_DIRS, MIN_FILE_SIZE_BYTES)\n\nif len(id_to_path) == 0:\n    print(\"No usable PNGs found — check INPUT_DIRS paths.\")\nelse:\n    df = load_labels(TRAIN_CSV)\n    df = df[df.index.isin(id_to_path.keys())]\n    print(f\"\\nMatched {len(df)} of {len(id_to_path)} scanned PNGs to labels.\")\n    print(f\"Positive rate ('any'): {df['any'].mean():.4f}\")\n    print(f\"\\nPer-subtype positive rates:\")\n    print(df[LABEL_COLS].mean().round(4))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:40:01.095097Z","iopub.execute_input":"2026-07-26T20:40:01.095424Z","iopub.status.idle":"2026-07-26T20:57:37.645067Z","shell.execute_reply.started":"2026-07-26T20:40:01.095384Z","shell.execute_reply":"2026-07-26T20:57:37.644313Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\ncomp_base = \"/kaggle/input/competitions/rsna-intracranial-hemorrhage-detection\"\n\n# just list top-level contents, don't recurse into everything\nprint(os.listdir(comp_base))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:57:37.646936Z","iopub.execute_input":"2026-07-26T20:57:37.647272Z","iopub.status.idle":"2026-07-26T20:57:37.652145Z","shell.execute_reply.started":"2026-07-26T20:57:37.647247Z","shell.execute_reply":"2026-07-26T20:57:37.651316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\npath = \"/kaggle/input/competitions/rsna-intracranial-hemorrhage-detection/rsna-intracranial-hemorrhage-detection\"\nprint(os.listdir(path))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:57:37.653296Z","iopub.execute_input":"2026-07-26T20:57:37.653634Z","iopub.status.idle":"2026-07-26T20:57:37.669167Z","shell.execute_reply.started":"2026-07-26T20:57:37.653591Z","shell.execute_reply":"2026-07-26T20:57:37.668292Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BalancedRandomSampler(Sampler):\n    def __init__(self, labels_any):\n        self.pos_idx = np.where(labels_any == 1)[0]\n        self.neg_idx = np.where(labels_any == 0)[0]\n\n    def __iter__(self):\n        n = min(len(self.neg_idx), max(len(self.pos_idx), 1))\n        neg_sample = np.random.choice(self.neg_idx, n, replace=False)\n        pos_sample = (np.random.choice(self.pos_idx, n, replace=False)\n                      if len(self.pos_idx) > 0 else np.array([], dtype=int))\n        ids = np.concatenate([pos_sample, neg_sample])\n        np.random.shuffle(ids)\n        return iter(ids.tolist())\n\n    def __len__(self):\n        return min(len(self.neg_idx), max(len(self.pos_idx), 1)) * 2\n\n\nclass ICHDataset(Dataset):\n    def __init__(self, df, id_to_path, transform=None):\n        self.df = df\n        self.id_to_path = id_to_path\n        self.transform = transform\n        self.ids = df.index.values\n        self.labels = df[LABEL_COLS].values.astype(np.float32)\n\n    def __len__(self):\n        return len(self.ids)\n\n    def __getitem__(self, idx):\n        img_id = self.ids[idx]\n        path = self.id_to_path[img_id]\n        img = Image.open(path).convert(\"RGB\")\n        if self.transform:\n            img = self.transform(img)\n        label = torch.tensor(self.labels[idx], dtype=torch.float32)\n        return img, label\n\n\ndef make_splits(df):\n    train_df, temp_df = train_test_split(\n        df, test_size=VAL_SPLIT + TEST_SPLIT, stratify=df[\"any\"], random_state=42\n    )\n    val_df, test_df = train_test_split(\n        temp_df, test_size=TEST_SPLIT / (VAL_SPLIT + TEST_SPLIT),\n        stratify=temp_df[\"any\"], random_state=42\n    )\n    return train_df, val_df, test_df\n\n\ntrain_df, val_df, test_df = make_splits(df)\nprint(f\"train: {len(train_df)}, val: {len(val_df)}, test: {len(test_df)}\")\nprint(f\"train positive rate: {train_df['any'].mean():.4f}\")\nprint(f\"val positive rate:   {val_df['any'].mean():.4f}\")\nprint(f\"test positive rate:  {test_df['any'].mean():.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:57:37.670082Z","iopub.execute_input":"2026-07-26T20:57:37.670578Z","iopub.status.idle":"2026-07-26T20:57:37.934478Z","shell.execute_reply.started":"2026-07-26T20:57:37.670555Z","shell.execute_reply":"2026-07-26T20:57:37.933769Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def build_vit_b16(num_classes=len(LABEL_COLS)):\n    model = models.vit_b_16(weights=models.ViT_B_16_Weights.IMAGENET1K_V1)\n    for param in model.parameters():\n        param.requires_grad = False\n    for param in model.encoder.layers[-1].parameters():\n        param.requires_grad = True\n    model.heads = nn.Sequential(\n        nn.Dropout(0.3),\n        nn.Linear(model.hidden_dim, num_classes),\n    )\n    for param in model.heads.parameters():\n        param.requires_grad = True\n    return model.to(DEVICE)\n\n\nmodel = build_vit_b16()\n\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\ntotal = sum(p.numel() for p in model.parameters())\nprint(f\"Trainable params: {trainable:,} / {total:,} ({100*trainable/total:.2f}%)\")\n\noptimizer = torch.optim.AdamW(\n    filter(lambda p: p.requires_grad, model.parameters()),\n    lr=1e-4, weight_decay=1e-4\n)\n\nwarmup = torch.optim.lr_scheduler.LinearLR(\n    optimizer, start_factor=0.1, end_factor=1.0, total_iters=1\n)\ncosine = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=max(NUM_EPOCHS - 1, 1)\n)\nscheduler = torch.optim.lr_scheduler.SequentialLR(\n    optimizer, schedulers=[warmup, cosine], milestones=[1]\n)\n\nprint(\"Model, optimizer, scheduler built successfully.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:57:37.93552Z","iopub.execute_input":"2026-07-26T20:57:37.935819Z","iopub.status.idle":"2026-07-26T20:57:41.386389Z","shell.execute_reply.started":"2026-07-26T20:57:37.935786Z","shell.execute_reply":"2026-07-26T20:57:41.385529Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ds = ICHDataset(train_df, id_to_path, transform=transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n]))\n\nval_ds = ICHDataset(val_df, id_to_path, transform=transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n]))\n\nsampler = BalancedRandomSampler(train_df[\"any\"].values)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                           num_workers=NUM_WORKERS, pin_memory=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                         num_workers=NUM_WORKERS, pin_memory=True)\n\nprint(f\"train batches: {len(train_loader)} | val batches: {len(val_loader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:57:41.387432Z","iopub.execute_input":"2026-07-26T20:57:41.387797Z","iopub.status.idle":"2026-07-26T20:57:41.403674Z","shell.execute_reply.started":"2026-07-26T20:57:41.387775Z","shell.execute_reply":"2026-07-26T20:57:41.402949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import time\n\ndef train_one_epoch(model, loader, optimizer, criterion, scaler, epoch):\n    model.train()\n    total_loss = 0.0\n    pbar = tqdm(loader, desc=f\"train epoch {epoch+1}/{NUM_EPOCHS}\", leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        optimizer.zero_grad()\n        with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        total_loss += loss.item() * imgs.size(0)\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n    return total_loss / len(loader.dataset)\n\n\n@torch.no_grad()\ndef evaluate(model, loader, criterion, epoch):\n    model.eval()\n    total_loss = 0.0\n    all_preds, all_labels = [], []\n    pbar = tqdm(loader, desc=f\"val epoch {epoch+1}/{NUM_EPOCHS}\", leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n        total_loss += loss.item() * imgs.size(0)\n        all_preds.append(torch.sigmoid(outputs.float()).cpu().numpy())\n        all_labels.append(labels.cpu().numpy())\n    all_preds = np.concatenate(all_preds)\n    all_labels = np.concatenate(all_labels)\n    aucs = {}\n    for i, name in enumerate(LABEL_COLS):\n        if len(np.unique(all_labels[:, i])) > 1:\n            aucs[name] = roc_auc_score(all_labels[:, i], all_preds[:, i])\n        else:\n            aucs[name] = float(\"nan\")\n    return total_loss / len(loader.dataset), aucs\n\n\ncriterion = nn.BCEWithLogitsLoss()\nscaler = torch.amp.GradScaler(device=DEVICE.type, enabled=USE_AMP)\nos.makedirs(CHECKPOINT_DIR, exist_ok=True)\ncheckpoint_path = os.path.join(CHECKPOINT_DIR, \"vit_b16_best.pt\")\n\nbest_val_loss = float(\"inf\")\nepochs_no_improve = 0\nstart = time.time()\n\nfor epoch in range(NUM_EPOCHS):\n    current_lr = optimizer.param_groups[0][\"lr\"]\n\n    train_loss = train_one_epoch(model, train_loader, optimizer, criterion, scaler, epoch)\n    val_loss, val_aucs = evaluate(model, val_loader, criterion, epoch)\n\n    auc_str = \", \".join(f\"{k}={v:.3f}\" for k, v in val_aucs.items())\n    elapsed = time.time() - start\n    print(f\"Epoch {epoch+1}/{NUM_EPOCHS} | lr={current_lr:.2e} | \"\n          f\"train_loss={train_loss:.4f} | val_loss={val_loss:.4f} | {auc_str} | \"\n          f\"elapsed={elapsed/60:.1f}min\")\n\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        epochs_no_improve = 0\n        torch.save({\n            \"epoch\": epoch + 1,\n            \"model_state_dict\": model.state_dict(),\n            \"optimizer_state_dict\": optimizer.state_dict(),\n            \"val_loss\": val_loss,\n            \"val_aucs\": val_aucs,\n        }, checkpoint_path)\n        print(f\"  → val loss improved, checkpoint saved to {checkpoint_path}\")\n    else:\n        epochs_no_improve += 1\n        print(f\"  → no improvement ({epochs_no_improve}/{EARLY_STOPPING_PATIENCE})\")\n\n    scheduler.step()\n\n    if epochs_no_improve >= EARLY_STOPPING_PATIENCE:\n        print(f\"Early stopping triggered at epoch {epoch+1}.\")\n        break\n\nprint(f\"\\nTraining complete. Best val_loss={best_val_loss:.4f}, checkpoint at {checkpoint_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T20:57:41.404676Z","iopub.execute_input":"2026-07-26T20:57:41.405164Z","iopub.status.idle":"2026-07-26T23:02:21.151069Z","shell.execute_reply.started":"2026-07-26T20:57:41.40514Z","shell.execute_reply":"2026-07-26T23:02:21.149928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchvision import transforms","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:02:21.162732Z","iopub.execute_input":"2026-07-26T23:02:21.163073Z","iopub.status.idle":"2026-07-26T23:02:21.168797Z","shell.execute_reply.started":"2026-07-26T23:02:21.163034Z","shell.execute_reply":"2026-07-26T23:02:21.168215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.metrics import (\n    roc_auc_score, average_precision_score, precision_recall_curve,\n    confusion_matrix, f1_score,\n)\n\n# Build test dataset/loader (val_ds/val_loader already exist from training)\neval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\ntest_ds = ICHDataset(test_df, id_to_path, transform=eval_transform)\ntest_loader = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=True)\n\n# Load best checkpoint\ncheckpoint = torch.load(checkpoint_path, map_location=DEVICE, weights_only=False)\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])\nmodel.eval()\nprint(f\"Loaded checkpoint: epoch {checkpoint['epoch']}, val_loss={checkpoint['val_loss']:.4f}\")\n\n\n@torch.no_grad()\ndef run_inference(loader):\n    all_probs, all_labels = [], []\n    for imgs, labels in loader:\n        imgs = imgs.to(DEVICE)\n        with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n            outputs = model(imgs)\n        probs = torch.sigmoid(outputs.float()).cpu().numpy()\n        all_probs.append(probs)\n        all_labels.append(labels.numpy())\n    return np.concatenate(all_probs), np.concatenate(all_labels)\n\n\nprint(\"Running inference on val split (threshold selection only)...\")\nval_probs, val_labels = run_inference(val_loader)\nprint(\"Running inference on test split (reported numbers)...\")\ntest_probs, test_labels = run_inference(test_loader)\nprint(f\"Done: {len(val_probs)} val images, {len(test_probs)} test images.\")\n\n\ndef best_f1_threshold(y_true, y_prob):\n    precisions, recalls, thresholds = precision_recall_curve(y_true, y_prob)\n    f1s = 2 * precisions * recalls / (precisions + recalls + 1e-12)\n    best_idx = np.nanargmax(f1s[:-1])\n    return thresholds[best_idx] if len(thresholds) > 0 else 0.5\n\n\nrows = []\nconfusion_matrices = {}\n\nfor i, name in enumerate(LABEL_COLS):\n    y_true_val = val_labels[:, i]\n    y_true_test = test_labels[:, i]\n    y_prob_test = test_probs[:, i]\n\n    if len(np.unique(y_true_val)) < 2 or len(np.unique(y_true_test)) < 2:\n        continue\n\n    thresh = best_f1_threshold(y_true_val, val_probs[:, i])\n    y_pred_test = (y_prob_test >= thresh).astype(int)\n\n    roc_auc = roc_auc_score(y_true_test, y_prob_test)\n    pr_auc = average_precision_score(y_true_test, y_prob_test)\n\n    tn, fp, fn, tp = confusion_matrix(y_true_test, y_pred_test, labels=[0, 1]).ravel()\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0.0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0.0\n    specificity = tn / (tn + fp) if (tn + fp) > 0 else 0.0\n    f1 = f1_score(y_true_test, y_pred_test, zero_division=0)\n\n    confusion_matrices[name] = np.array([[tn, fp], [fn, tp]])\n\n    rows.append({\n        \"class\": name, \"roc_auc\": roc_auc, \"pr_auc\": pr_auc,\n        \"threshold_from_val\": thresh, \"precision\": precision,\n        \"recall_sensitivity\": recall, \"specificity\": specificity, \"f1\": f1,\n        \"n_positive\": int(y_true_test.sum()), \"n_total\": len(y_true_test),\n    })\n\nresults_df = pd.DataFrame(rows).set_index(\"class\")\nmacro_auc = roc_auc_score(test_labels, test_probs, average=\"macro\")\nmicro_auc = roc_auc_score(test_labels, test_probs, average=\"micro\")\n\nprint(\"\\n\" + \"=\" * 90)\nprint(f\"PER-CLASS METRICS (test split) — vit_b16, checkpoint epoch {checkpoint['epoch']}\")\nprint(\"=\" * 90)\nprint(results_df.round(4).to_string())\nprint(f\"\\nMacro-average ROC-AUC: {macro_auc:.4f}\")\nprint(f\"Micro-average ROC-AUC: {micro_auc:.4f}\")\n\nprint(\"\\n\" + \"=\" * 90)\nprint(\"CONFUSION MATRICES (test split; rows=true, cols=predicted, order=[neg, pos])\")\nprint(\"=\" * 90)\nfor name, cm in confusion_matrices.items():\n    print(f\"\\n{name}:\")\n    print(f\"           pred_neg  pred_pos\")\n    print(f\"true_neg   {cm[0,0]:>8}  {cm[0,1]:>8}\")\n    print(f\"true_pos   {cm[1,0]:>8}  {cm[1,1]:>8}\")\n\nout_path = \"/kaggle/working/vit_b16_eval_metrics.csv\"\nresults_df.to_csv(out_path)\nprint(f\"\\nFull metrics table saved to {out_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:02:21.169762Z","iopub.execute_input":"2026-07-26T23:02:21.170102Z","iopub.status.idle":"2026-07-26T23:08:10.242356Z","shell.execute_reply.started":"2026-07-26T23:02:21.170078Z","shell.execute_reply":"2026-07-26T23:08:10.241013Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import brier_score_loss\n\nN_BINS = 15\n\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        return logits / torch.exp(self.log_T)\n\n@torch.no_grad()\ndef get_logits(loader):\n    all_logits, all_labels = [], []\n    for imgs, labels in loader:\n        imgs = imgs.to(DEVICE)\n        with torch.autocast(device_type=DEVICE.type, enabled=USE_AMP):\n            outputs = model(imgs)\n        all_logits.append(outputs.float().cpu())\n        all_labels.append(labels)\n    return torch.cat(all_logits), torch.cat(all_labels)\n\nprint(\"Running inference (logits) on val split (for fitting T)...\")\nval_logits, val_labels_t = get_logits(val_loader)\nprint(\"Running inference (logits) on test split (for reporting)...\")\ntest_logits, test_labels_t = get_logits(test_loader)\n\nscaler_t = TemperatureScaler()\noptimizer_t = torch.optim.LBFGS([scaler_t.log_T], lr=0.05, max_iter=100)\nbce = nn.BCEWithLogitsLoss()\n\ndef closure():\n    optimizer_t.zero_grad()\n    loss = bce(scaler_t(val_logits), val_labels_t)\n    loss.backward()\n    return loss\n\noptimizer_t.step(closure)\nT = torch.exp(scaler_t.log_T).item()\nprint(f\"\\nFitted temperature: T = {T:.4f} ({'softens' if T > 1 else 'sharpens'} confidence)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:08:10.244389Z","iopub.execute_input":"2026-07-26T23:08:10.244824Z","iopub.status.idle":"2026-07-26T23:12:54.713499Z","shell.execute_reply.started":"2026-07-26T23:08:10.244769Z","shell.execute_reply":"2026-07-26T23:12:54.712614Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def expected_and_max_calibration_error(y_true, y_prob, n_bins=15):\n    bin_edges = np.linspace(0, 1, n_bins + 1)\n    ece, mce = 0.0, 0.0\n    bin_accs, bin_confs, bin_counts = [], [], []\n    for lo, hi in zip(bin_edges[:-1], bin_edges[1:]):\n        mask = (y_prob > lo) & (y_prob <= hi)\n        count = mask.sum()\n        if count == 0:\n            bin_accs.append(np.nan)\n            bin_confs.append((lo + hi) / 2)\n            bin_counts.append(0)\n            continue\n        acc = y_true[mask].mean()\n        conf = y_prob[mask].mean()\n        gap = abs(acc - conf)\n        ece += (count / len(y_prob)) * gap\n        mce = max(mce, gap)\n        bin_accs.append(acc)\n        bin_confs.append(conf)\n        bin_counts.append(count)\n    return ece, mce, bin_edges, bin_accs, bin_confs, bin_counts\n\n\ndef nll(y_true, y_prob):\n    eps = 1e-12\n    p = np.clip(y_prob, eps, 1 - eps)\n    return -np.mean(y_true * np.log(p) + (1 - y_true) * np.log(1 - p))\n\n\ntest_labels_np = test_labels_t.numpy()\nprobs_before = torch.sigmoid(test_logits).numpy()\nprobs_after = torch.sigmoid(test_logits / T).numpy()\n\nrows = []\nfor i, name in enumerate(LABEL_COLS):\n    y_true = test_labels_np[:, i]\n    if len(np.unique(y_true)) < 2:\n        continue\n\n    p_before, p_after = probs_before[:, i], probs_after[:, i]\n    ece_b, mce_b, *_ = expected_and_max_calibration_error(y_true, p_before, N_BINS)\n    ece_a, mce_a, *_ = expected_and_max_calibration_error(y_true, p_after, N_BINS)\n\n    rows.append({\n        \"class\": name,\n        \"roc_auc\": roc_auc_score(y_true, p_before),\n        \"pr_auc\": average_precision_score(y_true, p_before),\n        \"ece_before\": ece_b, \"ece_after\": ece_a,\n        \"mce_before\": mce_b, \"mce_after\": mce_a,\n        \"brier_before\": brier_score_loss(y_true, p_before),\n        \"brier_after\": brier_score_loss(y_true, p_after),\n        \"nll_before\": nll(y_true, p_before),\n        \"nll_after\": nll(y_true, p_after),\n    })\n\ncalib_results_df = pd.DataFrame(rows).set_index(\"class\")\nprint(\"\\n\" + \"=\" * 100)\nprint(f\"CALIBRATION METRICS (test split, n={len(test_labels_np)}) — vit_b16, T={T:.4f}\")\nprint(\"=\" * 100)\nprint(calib_results_df.round(4).to_string())\n\nmacro_ece_before = calib_results_df[\"ece_before\"].mean()\nmacro_ece_after = calib_results_df[\"ece_after\"].mean()\nmacro_brier_before = calib_results_df[\"brier_before\"].mean()\nmacro_brier_after = calib_results_df[\"brier_after\"].mean()\nprint(f\"\\nMacro-avg ECE:   before={macro_ece_before:.4f} -> after={macro_ece_after:.4f}\")\nprint(f\"Macro-avg Brier: before={macro_brier_before:.4f} -> after={macro_brier_after:.4f}\")\n\nout_path = \"/kaggle/working/vit_b16_calibration_metrics.csv\"\ncalib_results_df.to_csv(out_path)\nwith open(\"/kaggle/working/vit_b16_temperature.txt\", \"w\") as f:\n    f.write(f\"T={T:.6f}\\n\")\nprint(f\"\\nSaved calibration table to {out_path}\")\nprint(f\"Saved fitted T to /kaggle/working/vit_b16_temperature.txt\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:12:54.714982Z","iopub.execute_input":"2026-07-26T23:12:54.715232Z","iopub.status.idle":"2026-07-26T23:12:54.90167Z","shell.execute_reply.started":"2026-07-26T23:12:54.715203Z","shell.execute_reply":"2026-07-26T23:12:54.901004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"any_idx = LABEL_COLS.index(\"any\")\ny_true_any = test_labels_np[:, any_idx]\np_before_any = probs_before[:, any_idx]\np_after_any = probs_after[:, any_idx]\n\nece_b, mce_b, edges_b, accs_b, confs_b, counts_b = expected_and_max_calibration_error(y_true_any, p_before_any, N_BINS)\nece_a, mce_a, edges_a, accs_a, confs_a, counts_a = expected_and_max_calibration_error(y_true_any, p_after_any, N_BINS)\n\nfig, axes = plt.subplots(2, 2, figsize=(11, 9))\nbin_centers = (edges_b[:-1] + edges_b[1:]) / 2\nbin_width = 1 / N_BINS\n\nfor col, (accs, counts, ece, mce, label) in enumerate([\n    (accs_b, counts_b, ece_b, mce_b, \"Before temp scaling (T=1.0)\"),\n    (accs_a, counts_a, ece_a, mce_a, f\"After temp scaling (T={T:.3f})\"),\n]):\n    valid = [not np.isnan(a) for a in accs]\n    ax_rel = axes[0, col]\n    ax_rel.bar(np.array(bin_centers)[valid], np.array(accs)[valid],\n               width=bin_width, edgecolor=\"black\", alpha=0.7, label=\"Model\")\n    ax_rel.plot([0, 1], [0, 1], \"k--\", label=\"Perfect calibration\")\n    ax_rel.set_xlabel(\"Predicted probability\")\n    ax_rel.set_ylabel(\"Observed frequency\")\n    ax_rel.set_title(f\"{label}\\nECE={ece:.4f}  MCE={mce:.4f}\")\n    ax_rel.legend(fontsize=8)\n\n    ax_hist = axes[1, col]\n    ax_hist.bar(bin_centers, counts, width=bin_width, edgecolor=\"black\", alpha=0.7, color=\"gray\")\n    ax_hist.set_xlabel(\"Predicted probability\")\n    ax_hist.set_ylabel(\"Count\")\n    ax_hist.set_title(\"Confidence histogram\")\n\nplt.suptitle(f\"Calibration — vit_b16, class 'any' (test, n={len(test_labels_np)})\", y=1.00)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/vit_b16_calibration_reliability.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"Saved reliability diagram to /kaggle/working/vit_b16_calibration_reliability.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:12:54.902746Z","iopub.execute_input":"2026-07-26T23:12:54.903128Z","iopub.status.idle":"2026-07-26T23:12:56.180955Z","shell.execute_reply.started":"2026-07-26T23:12:54.903103Z","shell.execute_reply":"2026-07-26T23:12:56.180049Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import roc_curve, precision_recall_curve, auc\n\nn_classes = len(LABEL_COLS)\nfig, axes = plt.subplots(2, n_classes, figsize=(4 * n_classes, 8))\n\nfor i, name in enumerate(LABEL_COLS):\n    y_true = test_labels[:, i]\n    y_prob = test_probs[:, i]\n    if len(np.unique(y_true)) < 2:\n        axes[0, i].set_visible(False)\n        axes[1, i].set_visible(False)\n        continue\n\n    # ROC\n    fpr, tpr, _ = roc_curve(y_true, y_prob)\n    roc_auc_val = auc(fpr, tpr)\n    ax_roc = axes[0, i]\n    ax_roc.plot(fpr, tpr, label=f\"AUC = {roc_auc_val:.3f}\")\n    ax_roc.plot([0, 1], [0, 1], \"k--\", alpha=0.4)\n    ax_roc.set_title(f\"{name}\\nROC curve\")\n    ax_roc.set_xlabel(\"False Positive Rate\")\n    ax_roc.set_ylabel(\"True Positive Rate\")\n    ax_roc.legend(fontsize=8, loc=\"lower right\")\n\n    # PR\n    precision, recall, _ = precision_recall_curve(y_true, y_prob)\n    pr_auc_val = average_precision_score(y_true, y_prob)\n    base_rate = y_true.mean()\n    ax_pr = axes[1, i]\n    ax_pr.plot(recall, precision, label=f\"AP = {pr_auc_val:.3f}\")\n    ax_pr.axhline(base_rate, color=\"k\", linestyle=\"--\", alpha=0.4, label=f\"baseline = {base_rate:.3f}\")\n    ax_pr.set_title(f\"{name}\\nPR curve\")\n    ax_pr.set_xlabel(\"Recall\")\n    ax_pr.set_ylabel(\"Precision\")\n    ax_pr.legend(fontsize=8, loc=\"upper right\")\n\nplt.suptitle(f\"ROC & PR curves — vit_b16, test split (n={len(test_labels)})\", y=1.02)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/vit_b16_roc_pr_curves.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\nprint(\"Saved ROC/PR curve grid to /kaggle/working/vit_b16_roc_pr_curves.png\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:12:56.182024Z","iopub.execute_input":"2026-07-26T23:12:56.182898Z","iopub.status.idle":"2026-07-26T23:12:59.136555Z","shell.execute_reply.started":"2026-07-26T23:12:56.182846Z","shell.execute_reply":"2026-07-26T23:12:59.135693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nprint(os.listdir('/kaggle/working'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:12:59.137707Z","iopub.execute_input":"2026-07-26T23:12:59.138059Z","iopub.status.idle":"2026-07-26T23:12:59.14249Z","shell.execute_reply.started":"2026-07-26T23:12:59.138033Z","shell.execute_reply":"2026-07-26T23:12:59.141797Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\n# Check current usage\nos.system(\"du -sh /kaggle/working/*\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:12:59.143453Z","iopub.execute_input":"2026-07-26T23:12:59.143758Z","iopub.status.idle":"2026-07-26T23:12:59.221543Z","shell.execute_reply.started":"2026-07-26T23:12:59.143736Z","shell.execute_reply":"2026-07-26T23:12:59.220763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfiles_to_delete = [\n    \"/kaggle/working/vit_outputs.zip\",\n    \"/kaggle/working/all_results.zip\",\n]\n\nfor f in files_to_delete:\n    if os.path.exists(f):\n        os.remove(f)\n        print(f\"Deleted: {f}\")\n    else:\n        print(f\"Not found: {f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:12:59.22249Z","iopub.execute_input":"2026-07-26T23:12:59.22278Z","iopub.status.idle":"2026-07-26T23:12:59.227641Z","shell.execute_reply.started":"2026-07-26T23:12:59.222748Z","shell.execute_reply":"2026-07-26T23:12:59.226928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.system(\"du -sh /kaggle/working/*\")\nos.system(\"df -h /kaggle/working\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:12:59.228606Z","iopub.execute_input":"2026-07-26T23:12:59.229132Z","iopub.status.idle":"2026-07-26T23:12:59.262285Z","shell.execute_reply.started":"2026-07-26T23:12:59.2291Z","shell.execute_reply":"2026-07-26T23:12:59.261495Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import FileLink\nFileLink(\"/kaggle/working/results_only_zip.zip\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-26T23:12:59.263213Z","iopub.execute_input":"2026-07-26T23:12:59.263527Z","iopub.status.idle":"2026-07-26T23:12:59.269153Z","shell.execute_reply.started":"2026-07-26T23:12:59.263503Z","shell.execute_reply":"2026-07-26T23:12:59.268226Z"}},"outputs":[],"execution_count":null}]}