{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"16e3c67a-0ead-401a-a879-bad8940c20fe","cell_type":"markdown","source":"# DRISHTI — APTOS 2019 Official Training Run 🎯\n\n**SIH 2026 | Problem Statement 26038 (MathWorks) | Team DRISHTI**\n\nThis notebook trains the **official DRISHTI DR classifier**:\n- **ResNet-50** (ImageNet-initialized) — matches our MATLAB implementation & PPT story\n- **5-class ICDR grading** (Levels 0–4) on the **APTOS 2019 dataset** (3,662 labeled images, Aravind Eye Hospital, rural India)\n- Reports the metrics the problem statement demands: **referable-DR sensitivity & specificity on a held-out test set (~550 images)**, plus AUC and Quadratic Weighted Kappa\n- Produces Grad-CAM explainability samples\n\n**How to run (teammate):**\n1. Attach the APTOS dataset: right panel → **Input** → **Add Input** → Competition → `aptos2019-blindness-detection`\n2. Settings (right panel): **Accelerator → GPU T4 x2**, **Internet → On**\n3. Click **Run All** and wait ~30–45 minutes\n4. When finished, download the output files (bottom of the right panel → Output) and send them to your mentor\n\n**Do not edit any code. If a cell fails, screenshot the FULL error and send it to the mentor.**","metadata":{}},{"id":"bad30282-f0fc-48d5-a7fa-6b0b4b86dfdb","cell_type":"code","source":"# ============================================================\n# CELL 1: setup, configuration, reproducibility\n# ============================================================\nimport os, sys, json, math, random, time, copy\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport matplotlib\nmatplotlib.use('Agg')\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.models as models\nimport torchvision.transforms as T\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix, cohen_kappa_score, roc_auc_score\n\ntry:\n    from tqdm.auto import tqdm\nexcept Exception:\n    def tqdm(x, **k):\n        return x\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('PyTorch :', torch.__version__)\nprint('Device  :', DEVICE)\nif DEVICE.type == 'cuda':\n    print('GPU     :', torch.cuda.get_device_name(0))\n    torch.backends.cudnn.benchmark = True\nelse:\n    print('WARNING: no GPU! Go to notebook Settings -> Accelerator -> GPU T4 x2, then restart and Run All again.')\n\n# ------------------- CONFIGURATION -------------------\nDATA_DIR    = '/kaggle/input/competitions/aptos2019-blindness-detection/'\nOUT_DIR     = '.'\nIMG_SIZE    = 256        # training image resolution\nBATCH       = 32\nEPOCHS      = 15\nLR          = 1e-4\nNUM_WORKERS = 2\nNUM_CLASSES = 5\nCLASS_NAMES = ['No DR (0)', 'Mild NPDR (1)', 'Moderate NPDR (2)', 'Severe NPDR (3)', 'Proliferative DR (4)']\nprint('Config OK: IMG_SIZE=%d  BATCH=%d  EPOCHS=%d' % (IMG_SIZE, BATCH, EPOCHS))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T23:36:58.311380Z","iopub.execute_input":"2026-08-30T23:36:58.311908Z","iopub.status.idle":"2026-08-30T23:36:58.323355Z","shell.execute_reply.started":"2026-08-30T23:36:58.311878Z","shell.execute_reply":"2026-08-30T23:36:58.322661Z"}},"outputs":[],"execution_count":null},{"id":"98862c24-10d0-4ca9-b0ec-41a7ceb2510c","cell_type":"code","source":"# ============================================================\n# CELL 2: load APTOS labels + stratified 70/15/15 split\n# (same split discipline as our STARE prototype)\n# ============================================================\ntrain_csv = os.path.join(DATA_DIR, 'train.csv')\nassert os.path.exists(train_csv), \\\n    'train.csv not found! Attach the APTOS competition data first (see instructions at the top).'\n\ndf = pd.read_csv(train_csv)\nprint('Total labeled images:', len(df))\nprint('Class distribution (doctor labels):')\nprint(df.diagnosis.value_counts().sort_index().to_string())\n\ntrain_df, tmp_df   = train_test_split(df,     test_size=0.30, stratify=df.diagnosis,     random_state=SEED)\nval_df,  test_df   = train_test_split(tmp_df, test_size=0.50, stratify=tmp_df.diagnosis, random_state=SEED)\n\nprint()\nprint('Split: train=%d  val=%d  test=%d   (70/15/15, stratified, seed 42)' % (len(train_df), len(val_df), len(test_df)))\nprint('Referable DR (level>=2) share: train %.1f%% | val %.1f%% | test %.1f%%' % (\n    100*(train_df.diagnosis >= 2).mean(),\n    100*(val_df.diagnosis   >= 2).mean(),\n    100*(test_df.diagnosis  >= 2).mean()))\nprint('NOTE: the test set is NEVER used for training or threshold tuning.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T23:37:02.199389Z","iopub.execute_input":"2026-08-30T23:37:02.199680Z","iopub.status.idle":"2026-08-30T23:37:02.219806Z","shell.execute_reply.started":"2026-08-30T23:37:02.199656Z","shell.execute_reply":"2026-08-30T23:37:02.219101Z"}},"outputs":[],"execution_count":null},{"id":"fd3495ea-e7b8-4a18-bd88-6ab191ea19bb","cell_type":"code","source":"# ============================================================\n# CELL 3: dataset with fundus cropping + augmentation\n# (augmentation recipe validated in our STARE prototype v3)\n# ============================================================\ndef crop_fundus(img):\n    \"\"\"Remove the black border around the retina (APTOS images vary in size).\"\"\"\n    if img is None:\n        return None\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    mask = gray > 12\n    if mask.sum() < 500:\n        return img\n    ys, xs = np.where(mask)\n    y1, y2 = max(0, ys.min() - 8), min(img.shape[0], ys.max() + 8)\n    x1, x2 = max(0, xs.min() - 8), min(img.shape[1], xs.max() + 8)\n    return img[y1:y2, x1:x2]\n\n\nclass FundusDataset(Dataset):\n    def __init__(self, df, img_dir, augment=False):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.augment = augment\n        self.tf_train = T.Compose([\n            T.ToPILImage(),\n            T.RandomHorizontalFlip(),\n            T.RandomVerticalFlip(p=0.3),\n            T.RandomRotation(20),\n            T.ColorJitter(brightness=0.25, contrast=0.25, saturation=0.2, hue=0.03),\n            T.RandomResizedCrop(IMG_SIZE, scale=(0.8, 1.0)),\n            T.ToTensor(),\n            T.RandomErasing(p=0.25, scale=(0.02, 0.08)),\n            T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n        ])\n        self.tf_val = T.Compose([\n            T.ToPILImage(),\n            T.ToTensor(),\n            T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n        ])\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        row = self.df.iloc[i]\n        path = os.path.join(self.img_dir, str(row.id_code) + '.png')\n        img = cv2.imread(path)\n        assert img is not None, 'cannot read image: ' + path\n        img = crop_fundus(img)\n        img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n        rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        x = self.tf_train(rgb) if self.augment else self.tf_val(rgb)\n        return x, int(row.diagnosis)\n\n\nIMG_DIR = os.path.join(DATA_DIR, 'train_images')\nassert os.path.exists(IMG_DIR), 'train_images folder not found! Attach the APTOS competition data first.'\n\nloaders = {\n    'train': DataLoader(FundusDataset(train_df, IMG_DIR, augment=True),\n                        batch_size=BATCH, shuffle=True,  num_workers=NUM_WORKERS,\n                        pin_memory=(DEVICE.type == 'cuda')),\n    'val':   DataLoader(FundusDataset(val_df,   IMG_DIR, augment=False),\n                        batch_size=BATCH, shuffle=False, num_workers=NUM_WORKERS,\n                        pin_memory=(DEVICE.type == 'cuda')),\n    'test':  DataLoader(FundusDataset(test_df,  IMG_DIR, augment=False),\n                        batch_size=BATCH, shuffle=False, num_workers=NUM_WORKERS,\n                        pin_memory=(DEVICE.type == 'cuda')),\n}\n\n# sanity check: push ONE image through the full pipeline\n_x, _y = FundusDataset(train_df, IMG_DIR, augment=False)[0]\nprint('Sanity check OK: image tensor', tuple(_x.shape), '| label', _y)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T23:37:06.129856Z","iopub.execute_input":"2026-08-30T23:37:06.130326Z","iopub.status.idle":"2026-08-30T23:37:06.454813Z","shell.execute_reply.started":"2026-08-30T23:37:06.130295Z","shell.execute_reply":"2026-08-30T23:37:06.453912Z"}},"outputs":[],"execution_count":null},{"id":"0c459fe1-76ee-47df-b20b-c87849900da6","cell_type":"code","source":"# ============================================================\n# CELL 4: the OFFICIAL DRISHTI model — ResNet-50, 5-class\n# + class-balanced loss (gentle sqrt weights + label smoothing,\n#   the recipe that fixed specificity in our STARE prototype v3)\n# ============================================================\ndef build_model():\n    net = None\n    try:\n        net = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)\n        print('Loaded ImageNet ResNet-50 weights (internet OK)')\n    except Exception as e:\n        print('Online weight download failed (%s) — trying offline copies...' % type(e).__name__)\n        net = models.resnet50()\n        loaded = False\n        for p in ['/kaggle/input/pretrained-pytorch-models/resnet50-19c8e357.pth',\n                  '/kaggle/input/pytorch-pretrained-image-models/resnet50-19c8e357.pth']:\n            if os.path.exists(p):\n                net.load_state_dict(torch.load(p, map_location='cpu'))\n                print('Loaded offline weights from:', p)\n                loaded = True\n                break\n        if not loaded:\n            print('WARNING: no pretrained weights found — training from scratch (results will be worse).')\n    net.fc = nn.Linear(net.fc.in_features, NUM_CLASSES)\n    return net.to(DEVICE)\n\n\nmodel = build_model()\n\ncounts = train_df.diagnosis.value_counts().sort_index().reindex(range(NUM_CLASSES)).fillna(1).values\nw = np.sqrt(len(train_df) / (NUM_CLASSES * counts))\nclass_weights = torch.tensor(w, dtype=torch.float32).to(DEVICE)\nprint('Train class counts :', counts.tolist())\nprint('Loss weights (sqrt):', np.round(w, 3).tolist())\n\ncriterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=0.1)\noptimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T23:37:12.287630Z","iopub.execute_input":"2026-08-30T23:37:12.288034Z","iopub.status.idle":"2026-08-30T23:37:13.757133Z","shell.execute_reply.started":"2026-08-30T23:37:12.288004Z","shell.execute_reply":"2026-08-30T23:37:13.755596Z"}},"outputs":[],"execution_count":null},{"id":"97807cd0-f7af-4216-b039-671dea2d160d","cell_type":"code","source":"# ============================================================\n# CELL 5: training loop (best model saved by validation QWK)\n# ============================================================\ndef evaluate(model, loader):\n    model.eval()\n    ys, ps, loss_sum, n = [], [], 0.0, 0\n    with torch.no_grad():\n        for x, y in loader:\n            x, y = x.to(DEVICE), y.to(DEVICE)\n            logits = model(x)\n            loss_sum += criterion(logits, y).item() * x.size(0)\n            n += x.size(0)\n            ps.append(torch.softmax(logits, 1).cpu())\n            ys.append(y.cpu())\n    P = torch.cat(ps).numpy()\n    Y = torch.cat(ys).numpy()\n    pred = P.argmax(1)\n    acc = float((pred == Y).mean())\n    qwk = float(cohen_kappa_score(Y, pred, weights='quadratic'))\n    ref_t, ref_p = (Y >= 2), (P[:, 2:].sum(1) >= 0.5)\n    sens = float(ref_p[ref_t].mean()) if ref_t.any() else 0.0\n    spec = float((~ref_p[~ref_t]).mean()) if (~ref_t).any() else 0.0\n    return loss_sum / n, acc, qwk, sens, spec, P, Y\n\n\nhistory, best_qwk, best_state = [], -1.0, None\nprint()\nprint('==================== TRAINING STARTS ====================')\nfor epoch in range(1, EPOCHS + 1):\n    t0 = time.time()\n    model.train()\n    tr_loss, tr_correct, tr_n = 0.0, 0, 0\n    for x, y in tqdm(loaders['train'], desc='epoch %d/%d' % (epoch, EPOCHS), leave=False):\n        x, y = x.to(DEVICE), y.to(DEVICE)\n        optimizer.zero_grad()\n        logits = model(x)\n        loss = criterion(logits, y)\n        loss.backward()\n        optimizer.step()\n        tr_loss += loss.item() * x.size(0)\n        tr_correct += (logits.argmax(1) == y).sum().item()\n        tr_n += x.size(0)\n    scheduler.step()\n\n    v_loss, v_acc, v_qwk, v_sens, v_spec, _, _ = evaluate(model, loaders['val'])\n    dt = time.time() - t0\n    history.append(dict(epoch=epoch, train_loss=tr_loss / tr_n, train_acc=tr_correct / tr_n,\n                        val_loss=v_loss, val_acc=v_acc, val_qwk=v_qwk,\n                        val_ref_sens=v_sens, val_ref_spec=v_spec))\n    star = ''\n    if v_qwk >= best_qwk:\n        best_qwk = v_qwk\n        best_state = copy.deepcopy(model.state_dict())\n        star = '   <- best so far (checkpoint saved)'\n    print('epoch %2d/%d | %3.0fs | train loss %.4f acc %.3f | VAL acc %.3f QWK %.3f | referable sens %.2f spec %.2f%s'\n          % (epoch, EPOCHS, dt, tr_loss / tr_n, tr_correct / tr_n, v_acc, v_qwk, v_sens, v_spec, star))\n\nprint()\nprint('Best validation QWK: %.4f' % best_qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-30T23:37:18.511912Z","iopub.execute_input":"2026-08-30T23:37:18.512324Z","iopub.status.idle":"2026-08-31T00:37:51.035612Z","shell.execute_reply.started":"2026-08-30T23:37:18.512294Z","shell.execute_reply":"2026-08-31T00:37:51.034168Z"}},"outputs":[],"execution_count":null},{"id":"8454cd1b-6ff3-4cc4-9583-71ac070d7f5b","cell_type":"code","source":"# ============================================================\n# CELL 6: operating-threshold selection (on VALIDATION only!)\n#          + FINAL evaluation on the HELD-OUT TEST SET\n# ============================================================\nmodel.load_state_dict(best_state)\n_, _, _, _, _, P_val,  Y_val  = evaluate(model, loaders['val'])\n_, _, _, _, _, P_test, Y_test = evaluate(model, loaders['test'])\n\np_ref_val  = P_val[:, 2:].sum(1)    # P(level >= 2) = referable DR probability\np_ref_test = P_test[:, 2:].sum(1)\n\n# ---- choose the referable-DR threshold on VALIDATION (never on test) ----\nrows = []\nfor thr in np.arange(0.05, 0.96, 0.01):\n    sens_v = float((p_ref_val >= thr)[Y_val >= 2].mean())\n    spec_v = float((p_ref_val <  thr)[Y_val <  2].mean())\n    rows.append((float(thr), sens_v, spec_v))\neligible = [r for r in rows if r[1] >= 0.90]\nif eligible:\n    THR = max(eligible, key=lambda r: r[2])[0]\n    thr_note = 'validation sensitivity >= 90% achieved; threshold set to maximize specificity'\nelse:\n    THR = max(rows, key=lambda r: r[1] + r[2])[0]\n    thr_note = 'WARNING: 90% sensitivity not reachable on validation — threshold set to best balanced point. Report this honestly!'\nprint('Chosen referable-DR threshold: %.2f' % THR)\nprint('(%s)' % thr_note)\n\n# ---- FINAL TEST METRICS ----\npred5 = P_test.argmax(1)\nref_true, ref_pred = (Y_test >= 2), (p_ref_test >= THR)\ntp = int((ref_true &  ref_pred).sum())\nfn = int((ref_true & ~ref_pred).sum())\ntn = int((~ref_true & ~ref_pred).sum())\nfp = int((~ref_true &  ref_pred).sum())\nsens = tp / (tp + fn)\nspec = tn / (tn + fp)\nacc  = float((pred5 == Y_test).mean())\nqwk  = float(cohen_kappa_score(Y_test, pred5, weights='quadratic'))\n\ntry:\n    auc_macro = float(roc_auc_score(Y_test, P_test, multi_class='ovr', average='macro'))\nexcept Exception as e:\n    auc_macro = None\n    print('(macro AUC skipped: %s)' % type(e).__name__)\ntry:\n    auc_ref = float(roc_auc_score(ref_true.astype(int), p_ref_test))\nexcept Exception as e:\n    auc_ref = None\n    print('(referable AUC skipped: %s)' % type(e).__name__)\n\ncm = confusion_matrix(Y_test, pred5, labels=list(range(NUM_CLASSES)))\nper_class_recall = {}\nfor i in range(NUM_CLASSES):\n    per_class_recall[CLASS_NAMES[i]] = round(float(cm[i, i] / cm[i].sum()), 3) if cm[i].sum() else None\n\nprint()\nprint('=' * 70)\nprint('FINAL HELD-OUT TEST RESULTS  (n = %d images, never seen during training)' % len(Y_test))\nprint('=' * 70)\nprint('REFERABLE DR (level >= 2) at threshold %.2f:' % THR)\nprint('   SENSITIVITY : %5.1f%%   (found %d of %d true DR cases)' % (100 * sens, tp, tp + fn))\nprint('   SPECIFICITY : %5.1f%%   (correctly cleared %d of %d non-referable)' % (100 * spec, tn, tn + fp))\nprint('   false referrals: %d  |  missed DR cases: %d' % (fp, fn))\nprint('Overall 5-class accuracy : %.1f%%' % (100 * acc))\nprint('Quadratic Weighted Kappa : %.3f' % qwk)\nif auc_macro:\n    print('AUC (macro, one-vs-rest)  : %.3f' % auc_macro)\nif auc_ref:\n    print('AUC (referable DR)        : %.3f' % auc_ref)\nprint('Per-class recall:', json.dumps(per_class_recall))\n\nmet_s, met_p = sens > 0.90, spec > 0.85\nif met_s and met_p:\n    verdict = 'BOTH PS TARGETS MET (sens>90%, spec>85%)'\nelif met_s:\n    verdict = 'sensitivity target met; specificity below 85% — report honestly'\nelif met_p:\n    verdict = 'specificity target met; sensitivity below 90% — report honestly'\nelse:\n    verdict = 'targets not met on this run — report honestly'\nprint()\nprint('PROBLEM-STATEMENT TARGETS (sensitivity>90%, specificity>85%):')\nprint('   -> ' + verdict)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-31T00:38:27.308033Z","iopub.execute_input":"2026-08-31T00:38:27.308557Z","iopub.status.idle":"2026-08-31T00:39:51.124338Z","shell.execute_reply.started":"2026-08-31T00:38:27.308461Z","shell.execute_reply":"2026-08-31T00:39:51.123458Z"}},"outputs":[],"execution_count":null},{"id":"43de8d46-2c7f-4079-ba90-4a3059106e0d","cell_type":"code","source":"# ============================================================\n# CELL 7: training curves + confusion matrix (for your slides!)\n# ============================================================\nfig, axes = plt.subplots(1, 3, figsize=(16, 4.2))\nep = [h['epoch'] for h in history]\naxes[0].plot(ep, [h['train_loss'] for h in history], 'o-', label='train loss')\naxes[0].plot(ep, [h['val_loss'] for h in history], 's-', label='val loss')\naxes[0].set_title('Loss'); axes[0].set_xlabel('epoch'); axes[0].legend(); axes[0].grid(alpha=0.3)\naxes[1].plot(ep, [h['train_acc'] for h in history], 'o-', label='train acc')\naxes[1].plot(ep, [h['val_acc'] for h in history], 's-', label='val acc')\naxes[1].set_title('Accuracy'); axes[1].set_xlabel('epoch'); axes[1].legend(); axes[1].grid(alpha=0.3)\naxes[2].plot(ep, [h['val_qwk'] for h in history], 'o-', color='green', label='val QWK')\naxes[2].plot(ep, [h['val_ref_sens'] for h in history], 's--', color='red', label='val referable sens')\naxes[2].plot(ep, [h['val_ref_spec'] for h in history], '^--', color='blue', label='val referable spec')\naxes[2].set_title('Validation metrics'); axes[2].set_xlabel('epoch'); axes[2].legend(); axes[2].grid(alpha=0.3)\nplt.tight_layout()\nplt.savefig(os.path.join(OUT_DIR, 'drishti_aptos_training_curves.png'), dpi=120)\nplt.show()\n\nfig, ax = plt.subplots(figsize=(7, 6))\nax.imshow(cm, cmap='Blues')\nax.set_xticks(range(NUM_CLASSES)); ax.set_yticks(range(NUM_CLASSES))\nax.set_xticklabels([c.split(' (')[0] for c in CLASS_NAMES], rotation=30, ha='right')\nax.set_yticklabels([c.split(' (')[0] for c in CLASS_NAMES])\nax.set_xlabel('AI prediction'); ax.set_ylabel(\"Doctor's label\")\nax.set_title('DRISHTI (ResNet-50) on APTOS — test confusion matrix (n=%d)' % len(Y_test))\nfor i in range(NUM_CLASSES):\n    for j in range(NUM_CLASSES):\n        ax.text(j, i, cm[i, j], ha='center', va='center', fontsize=11,\n                color='white' if cm[i, j] > cm.max() / 2 else 'black')\nplt.tight_layout()\nplt.savefig(os.path.join(OUT_DIR, 'drishti_aptos_confusion_matrix.png'), dpi=120)\nplt.show()\nprint('Saved: drishti_aptos_training_curves.png + drishti_aptos_confusion_matrix.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-31T00:43:43.316116Z","iopub.execute_input":"2026-08-31T00:43:43.316423Z","iopub.status.idle":"2026-08-31T00:43:44.092990Z","shell.execute_reply.started":"2026-08-31T00:43:43.316386Z","shell.execute_reply":"2026-08-31T00:43:44.092082Z"}},"outputs":[],"execution_count":null},{"id":"4d0625e6-e6aa-42e5-82bc-b0d092061ffa","cell_type":"code","source":"# ============================================================\n# CELL 8: Grad-CAM samples (explainability proof, like Module 4)\n# ============================================================\ndef gradcam_for(model, img_bgr):\n    img = crop_fundus(img_bgr)\n    img = cv2.resize(img, (IMG_SIZE, IMG_SIZE))\n    rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    tf = T.Compose([T.ToPILImage(), T.ToTensor(),\n                    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])])\n    x = tf(rgb).unsqueeze(0).requires_grad_(True).to(DEVICE)\n    acts, grads = {}, {}\n    h1 = model.layer4[-1].conv2.register_forward_hook(\n        lambda m, i, o: acts.__setitem__('a', o.detach()))\n    h2 = model.layer4[-1].conv2.register_full_backward_hook(\n        lambda m, gi, go: grads.__setitem__('g', go[0].detach()))\n    try:\n        logits = model(x)\n        cls = int(logits.argmax(1))\n        logits[0, cls].backward()\n    finally:\n        h1.remove(); h2.remove()\n    w_ = grads['g'].mean(dim=(2, 3), keepdim=True)\n    cam = F.relu((w_ * acts['a']).sum(1))[0].cpu()\n    cam = F.interpolate(cam[None, None], size=(img.shape[0], img.shape[1]),\n                        mode='bilinear', align_corners=False)[0, 0].numpy()\n    cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)\n    return cls, cam, rgb\n\n\nidx_ref = np.where(Y_test >= 2)[0][:2].tolist()\nidx_no  = np.where(Y_test <  2)[0][:2].tolist()\nfig, axes = plt.subplots(2, 4, figsize=(16, 8))\nfor row, idxs in enumerate([idx_ref, idx_no]):\n    for col, ti in enumerate(idxs[:2]):\n        code = str(test_df.iloc[ti].id_code)\n        img = cv2.imread(os.path.join(IMG_DIR, code + '.png'))\n        cls, cam, rgb = gradcam_for(model, img)\n        heat = cv2.applyColorMap((cam * 255).astype(np.uint8), cv2.COLORMAP_JET)\n        heat = cv2.cvtColor(heat, cv2.COLOR_BGR2RGB)\n        overlay = (0.45 * rgb + 0.55 * heat).astype(np.uint8)\n        axes[row, 2 * col].imshow(rgb); axes[row, 2 * col].axis('off')\n        axes[row, 2 * col].set_title('true: %s | AI: %s' % (\n            CLASS_NAMES[Y_test[ti]].split(' (')[0], CLASS_NAMES[cls].split(' (')[0]), fontsize=10)\n        axes[row, 2 * col + 1].imshow(overlay); axes[row, 2 * col + 1].axis('off')\n        axes[row, 2 * col + 1].set_title('Grad-CAM attention', fontsize=10)\nplt.suptitle('DRISHTI explainability on APTOS test images (top rows: referable DR, bottom: non-referable)', fontsize=12)\nplt.tight_layout()\nplt.savefig(os.path.join(OUT_DIR, 'drishti_aptos_gradcam_samples.png'), dpi=120)\nplt.show()\nprint('Saved: drishti_aptos_gradcam_samples.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-31T00:43:51.483425Z","iopub.execute_input":"2026-08-31T00:43:51.483977Z","iopub.status.idle":"2026-08-31T00:43:55.373247Z","shell.execute_reply.started":"2026-08-31T00:43:51.483945Z","shell.execute_reply":"2026-08-31T00:43:55.372424Z"}},"outputs":[],"execution_count":null},{"id":"9274d97e-8ced-4c5b-9ebd-edbb3db85a2e","cell_type":"code","source":"# ============================================================\n# CELL 9: save EVERYTHING (download these files afterwards!)\n# ============================================================\nmetrics = {\n    'dataset': 'APTOS 2019 (train split 70/15/15, stratified, seed 42)',\n    'model': 'ResNet-50 (ImageNet init), 5-class ICDR — official DRISHTI architecture',\n    'n_train': int(len(train_df)), 'n_val': int(len(val_df)), 'n_test': int(len(Y_test)),\n    'referable_threshold': round(float(THR), 2),\n    'threshold_note': thr_note,\n    'referable_sensitivity': round(float(sens), 4),\n    'referable_specificity': round(float(spec), 4),\n    'false_referrals': fp, 'missed_dr_cases': fn,\n    'accuracy_5class': round(float(acc), 4),\n    'quadratic_weighted_kappa': round(qwk, 4),\n    'auc_macro_ovr': (round(auc_macro, 4) if auc_macro else None),\n    'auc_referable': (round(auc_ref, 4) if auc_ref else None),\n    'per_class_recall': per_class_recall,\n    'confusion_matrix': cm.tolist(),\n    'class_names': CLASS_NAMES,\n    'img_size': IMG_SIZE,\n    'epochs': EPOCHS,\n    'best_val_qwk': round(float(best_qwk), 4),\n}\nwith open(os.path.join(OUT_DIR, 'drishti_aptos_results.json'), 'w') as f:\n    json.dump(metrics, f, indent=2)\n\ntorch.save({'state_dict': best_state, 'classes': CLASS_NAMES,\n            'img_size': IMG_SIZE, 'referable_threshold': round(float(THR), 2)},\n           os.path.join(OUT_DIR, 'drishti_aptos_resnet50.pt'))\n\npreds_out = test_df[['id_code', 'diagnosis']].copy()\npreds_out['predicted_level'] = pred5\npreds_out['p_referable'] = np.round(p_ref_test, 4)\npreds_out.to_csv(os.path.join(OUT_DIR, 'drishti_aptos_test_predictions.csv'), index=False)\n\nprint('=' * 70)\nprint('ALL DONE! Files saved in this notebook output (download them all):')\nfor f in ['drishti_aptos_results.json',\n          'drishti_aptos_resnet50.pt',\n          'drishti_aptos_confusion_matrix.png',\n          'drishti_aptos_training_curves.png',\n          'drishti_aptos_gradcam_samples.png',\n          'drishti_aptos_test_predictions.csv']:\n    print('   - ' + f)\nprint()\nprint('NEXT STEP: download these files and send them to your mentor.')\nprint('Screenshot the FINAL TEST RESULTS block above too — that is your new Slide 13!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-08-31T00:44:02.820321Z","iopub.execute_input":"2026-08-31T00:44:02.820905Z","iopub.status.idle":"2026-08-31T00:44:03.031369Z","shell.execute_reply.started":"2026-08-31T00:44:02.820873Z","shell.execute_reply":"2026-08-31T00:44:03.030427Z"}},"outputs":[],"execution_count":null},{"id":"7e0b4d47-537b-4a45-8255-3b339155d1a3","cell_type":"markdown","source":"## 📤 After the run finishes — what to do\n\n1. **Right panel → Output → download all 6 files**\n2. Send to your mentor (chat or USB):\n   - `drishti_aptos_results.json` ← the real metrics\n   - `drishti_aptos_confusion_matrix.png` ← goes on Slide 13\n   - `drishti_aptos_training_curves.png` ← proof of honest training\n   - `drishti_aptos_gradcam_samples.png` ← explainability proof\n   - a **screenshot of the FINAL TEST RESULTS** printout\n3. Keep `drishti_aptos_resnet50.pt` and `drishti_aptos_test_predictions.csv` safe —\n   we need them for tomorrow's **\"integrated pipeline vs single CNN\"** experiment.\n\n**Honesty rule:** whatever numbers come out — those are our numbers. We report them exactly as they are.","metadata":{}}]}