{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042},{"sourceType":"datasetVersion","sourceId":14968571,"datasetId":8562918,"databundleVersionId":15840513}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Set-Up**","metadata":{}},{"cell_type":"code","source":"%pip install torchio --q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:04:25.670996Z","iopub.execute_input":"2026-02-26T21:04:25.671184Z","iopub.status.idle":"2026-02-26T21:04:30.906758Z","shell.execute_reply.started":"2026-02-26T21:04:25.671164Z","shell.execute_reply":"2026-02-26T21:04:30.905812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#@title Imports\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport pytorch_lightning as pl\nimport torchio as tio \nimport torchvision\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader\n\nimport pydicom\n\nimport math\nimport os\nimport random\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom IPython.display import clear_output\nfrom tqdm.notebook import trange, tqdm\nfrom pathlib import Path\nfrom tqdm import tqdm\n\n%matplotlib inline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:04:30.909094Z","iopub.execute_input":"2026-02-26T21:04:30.909422Z","iopub.status.idle":"2026-02-26T21:04:46.131735Z","shell.execute_reply.started":"2026-02-26T21:04:30.909392Z","shell.execute_reply":"2026-02-26T21:04:46.131157Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Pre-Processing**","metadata":{}},{"cell_type":"code","source":"labels_df = pd.read_csv(\n    \"/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_train_labels.csv\"\n)\n\nlabels_df = labels_df.groupby(\"patientId\")[\"Target\"].max().reset_index()\n\nprint(labels_df[\"Target\"].value_counts())\nprint(labels_df.info(verbose=True, show_counts=True))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:04:46.132657Z","iopub.execute_input":"2026-02-26T21:04:46.133320Z","iopub.status.idle":"2026-02-26T21:04:46.222130Z","shell.execute_reply.started":"2026-02-26T21:04:46.133274Z","shell.execute_reply":"2026-02-26T21:04:46.221496Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"labels_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:04:46.223038Z","iopub.execute_input":"2026-02-26T21:04:46.223357Z","iopub.status.idle":"2026-02-26T21:04:46.241699Z","shell.execute_reply.started":"2026-02-26T21:04:46.223328Z","shell.execute_reply":"2026-02-26T21:04:46.241146Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Subjects","metadata":{}},{"cell_type":"code","source":"ROOT_PATH = Path(\"/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_train_images/\")\npatient_dirs = list(ROOT_PATH.glob(\"*\"))\n\npatient_dirs[0] # debug","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:04:46.242644Z","iopub.execute_input":"2026-02-26T21:04:46.242962Z","iopub.status.idle":"2026-02-26T21:04:47.056774Z","shell.execute_reply.started":"2026-02-26T21:04:46.242927Z","shell.execute_reply":"2026-02-26T21:04:47.055985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ds = pydicom.dcmread(patient_dirs[0])\nimg = ds.pixel_array\n\nimg.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:04:47.057863Z","iopub.execute_input":"2026-02-26T21:04:47.058358Z","iopub.status.idle":"2026-02-26T21:04:47.119265Z","shell.execute_reply.started":"2026-02-26T21:04:47.058331Z","shell.execute_reply":"2026-02-26T21:04:47.118694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_patient_label(patient_path: Path, labels_df: pd.DataFrame) -> int:\n    patientID = patient_path.stem\n    label = labels_df.loc[labels_df[\"patientId\"] == patientID, \"Target\"]\n    label = label.iloc[0] if not label.empty else None\n    return int(label)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:04:47.121534Z","iopub.execute_input":"2026-02-26T21:04:47.122086Z","iopub.status.idle":"2026-02-26T21:04:47.125950Z","shell.execute_reply.started":"2026-02-26T21:04:47.122061Z","shell.execute_reply":"2026-02-26T21:04:47.125180Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"subjects = []\nheights = []\nwidths = []\nlabels = []\n\nfor subject_path in tqdm(patient_dirs):\n    \n    img_path = subject_path\n    label = get_patient_label(subject_path, labels_df)\n\n    ct = tio.ScalarImage(img_path)\n    h, w, _ = ct.spatial_shape   \n\n    subject = tio.Subject(\n        CT = ct,\n        Label = torch.tensor(label, dtype=torch.long),\n        PatientID = subject_path.stem\n    )\n\n    subjects.append(subject)\n    heights.append(h)\n    widths.append(w)\n    labels.append(label)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:04:47.127022Z","iopub.execute_input":"2026-02-26T21:04:47.127299Z","iopub.status.idle":"2026-02-26T21:11:25.400887Z","shell.execute_reply.started":"2026-02-26T21:04:47.127265Z","shell.execute_reply":"2026-02-26T21:11:25.399980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(type(subjects[15][\"CT\"]), subjects[15][\"CT\"])\nprint(type(subjects[15][\"Label\"]), subjects[15][\"Label\"])\nsubjects[15][\"CT\"].spatial_shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:25.402097Z","iopub.execute_input":"2026-02-26T21:11:25.402447Z","iopub.status.idle":"2026-02-26T21:11:25.444550Z","shell.execute_reply.started":"2026-02-26T21:11:25.402419Z","shell.execute_reply":"2026-02-26T21:11:25.443969Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot depth (dimensions)\nplt.figure(figsize=(18,5))\n\n# Plot height\nplt.subplot(1,2,1)\nplt.hist(heights, bins=20, color='lightgreen', edgecolor='black')\nplt.title(\"CT Height Distribution\")\nplt.xlabel(\"Height (pixels/voxels)\")\nplt.ylabel(\"Number of Subjects\")\n\n# Plot width\nplt.subplot(1,2,2)\nplt.hist(widths, bins=20, color='salmon', edgecolor='black')\nplt.title(\"CT Width Distribution\")\nplt.xlabel(\"Width (pixels/voxels)\")\nplt.ylabel(\"Number of Subjects\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:25.445398Z","iopub.execute_input":"2026-02-26T21:11:25.445723Z","iopub.status.idle":"2026-02-26T21:11:25.896506Z","shell.execute_reply.started":"2026-02-26T21:11:25.445673Z","shell.execute_reply":"2026-02-26T21:11:25.895899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_size_og = subjects[15][\"CT\"].spatial_shape[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:25.897445Z","iopub.execute_input":"2026-02-26T21:11:25.897795Z","iopub.status.idle":"2026-02-26T21:11:25.906927Z","shell.execute_reply.started":"2026-02-26T21:11:25.897741Z","shell.execute_reply":"2026-02-26T21:11:25.906010Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **Transforms**","metadata":{}},{"cell_type":"code","source":"process = tio.Compose([\n    tio.ToCanonical(),                                 # step 1: fix orientation - RAS              \n    tio.RescaleIntensity((0, 1)),                      # step 2: normalize intensity\n    tio.Resize((356, 356, 1)),\n    tio.CropOrPad((256, 256, 1)),          \n])\n\naugmentation = tio.RandomAffine(scales=(0.9, 1.1), degrees=(-10, 10))\n\ntrain_transform = tio.Compose([process, augmentation])\nval_transform = tio.Compose([process])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:25.907973Z","iopub.execute_input":"2026-02-26T21:11:25.908257Z","iopub.status.idle":"2026-02-26T21:11:25.913456Z","shell.execute_reply.started":"2026-02-26T21:11:25.908217Z","shell.execute_reply":"2026-02-26T21:11:25.912744Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **DataSet & DataLoader**","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\n# (90/10)\ntrain_val_subjects, test_subjects = train_test_split(\n    subjects,\n    test_size=0.15,\n    stratify=labels,\n    random_state=42\n)\n\n# (80/20)\ntrain_subjects, val_subjects = train_test_split(\n    train_val_subjects,\n    test_size=0.2,\n    stratify=[s.Label.item() for s in train_val_subjects],\n    random_state=42\n)\n\n# Verify class distributions\ntrain_labels = [s.Label.item() for s in train_subjects]\nval_labels   = [s.Label.item() for s in val_subjects]\ntest_labels  = [s.Label.item() for s in test_subjects]\n\nprint(\"Train counts:\", np.bincount(train_labels))\nprint(\"Val counts:\", np.bincount(val_labels))\nprint(\"Test counts:\", np.bincount(test_labels))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:25.914272Z","iopub.execute_input":"2026-02-26T21:11:25.914518Z","iopub.status.idle":"2026-02-26T21:11:26.031535Z","shell.execute_reply.started":"2026-02-26T21:11:25.914495Z","shell.execute_reply":"2026-02-26T21:11:26.030939Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = tio.SubjectsDataset(train_subjects, transform = train_transform) \nval_dataset = tio.SubjectsDataset(val_subjects, transform = val_transform)  \ntest_dataset = tio.SubjectsDataset(test_subjects, transform = val_transform)  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:26.032393Z","iopub.execute_input":"2026-02-26T21:11:26.032974Z","iopub.status.idle":"2026-02-26T21:11:26.038948Z","shell.execute_reply.started":"2026-02-26T21:11:26.032946Z","shell.execute_reply":"2026-02-26T21:11:26.038292Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from typing import Tuple, List\n\ndef collate_subjects(batch: List) -> Tuple[torch.Tensor, torch.Tensor]:\n    images = torch.stack([s.CT.data.squeeze(-1).transpose(-2, -1).flip(1) for s in batch]) # transose to set orientation up right\n    labels = torch.tensor([s.Label.item() for s in batch])\n    return images, labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:26.039811Z","iopub.execute_input":"2026-02-26T21:11:26.040118Z","iopub.status.idle":"2026-02-26T21:11:26.052974Z","shell.execute_reply.started":"2026-02-26T21:11:26.040091Z","shell.execute_reply":"2026-02-26T21:11:26.052306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(train_dataset, batch_size=128, num_workers=4, collate_fn=collate_subjects, shuffle=True, pin_memory=True)\nval_loader = DataLoader(val_dataset, batch_size=128, num_workers=4, collate_fn=collate_subjects)\ntest_loader = DataLoader(test_dataset, batch_size=128, num_workers=4, collate_fn=collate_subjects)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:26.053949Z","iopub.execute_input":"2026-02-26T21:11:26.054243Z","iopub.status.idle":"2026-02-26T21:11:26.064142Z","shell.execute_reply.started":"2026-02-26T21:11:26.054218Z","shell.execute_reply":"2026-02-26T21:11:26.063516Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(train_loader))\nprint(f\"Images shape fresh off the loader: {x.shape}\")\nprint(f\"Labels shape fresh off the loader: {y.shape}\")\nprint(f\"Labels corresponding to {y.shape[0]} images in the batch: \" + str(y))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:26.064881Z","iopub.execute_input":"2026-02-26T21:11:26.065141Z","iopub.status.idle":"2026-02-26T21:11:49.465547Z","shell.execute_reply.started":"2026-02-26T21:11:26.065102Z","shell.execute_reply":"2026-02-26T21:11:49.464566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def imshow(img):\n    img = img.squeeze(0)\n    plt.imshow(img, cmap='gray')\n    plt.axis(\"off\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:49.467590Z","iopub.execute_input":"2026-02-26T21:11:49.468045Z","iopub.status.idle":"2026-02-26T21:11:49.474066Z","shell.execute_reply.started":"2026-02-26T21:11:49.467975Z","shell.execute_reply":"2026-02-26T21:11:49.473175Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_images = 16\ncols = 4\nrows = n_images // cols\n\nplt.figure(figsize=(12, 12))\n\nfor i in range(n_images):\n    plt.subplot(rows, cols, i + 1)\n    imshow(x[i])\n    plt.title(f\"Label: {y[i].item()}\")\n    plt.axis(\"off\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:49.475293Z","iopub.execute_input":"2026-02-26T21:11:49.475693Z","iopub.status.idle":"2026-02-26T21:11:50.825800Z","shell.execute_reply.started":"2026-02-26T21:11:49.475649Z","shell.execute_reply":"2026-02-26T21:11:50.825035Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Model**","metadata":{}},{"cell_type":"code","source":"img_size = x.shape[2]\nprint(f\"Image size fresh off the loader: {img_size}\")\n\nchannels_in = x.shape[1]\nprint(f\"Input Channels: {channels_in}\")\n\nnum_classes = 1\nprint(f\"Output Channels: {num_classes}\")\n\nlearning_rate = 3e-4","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:50.829081Z","iopub.execute_input":"2026-02-26T21:11:50.829315Z","iopub.status.idle":"2026-02-26T21:11:50.834145Z","shell.execute_reply.started":"2026-02-26T21:11:50.829289Z","shell.execute_reply":"2026-02-26T21:11:50.833337Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision.models as models\nmodel = models.resnet50(pretrained=True);\nmodel;","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:50.837781Z","iopub.execute_input":"2026-02-26T21:11:50.838142Z","iopub.status.idle":"2026-02-26T21:11:51.897735Z","shell.execute_reply.started":"2026-02-26T21:11:50.838118Z","shell.execute_reply":"2026-02-26T21:11:51.896990Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.conv1 = torch.nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n\nnum_features = model.fc.in_features\nmodel.fc = nn.Linear(num_features, 1)\n\nmodel;","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:51.898471Z","iopub.execute_input":"2026-02-26T21:11:51.898669Z","iopub.status.idle":"2026-02-26T21:11:51.904499Z","shell.execute_reply.started":"2026-02-26T21:11:51.898647Z","shell.execute_reply":"2026-02-26T21:11:51.903662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:51.905626Z","iopub.execute_input":"2026-02-26T21:11:51.905914Z","iopub.status.idle":"2026-02-26T21:11:51.917183Z","shell.execute_reply.started":"2026-02-26T21:11:51.905887Z","shell.execute_reply":"2026-02-26T21:11:51.916554Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.to(device);","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:51.918151Z","iopub.execute_input":"2026-02-26T21:11:51.918852Z","iopub.status.idle":"2026-02-26T21:11:51.974046Z","shell.execute_reply.started":"2026-02-26T21:11:51.918795Z","shell.execute_reply":"2026-02-26T21:11:51.973387Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Train**","metadata":{}},{"cell_type":"code","source":"optimizer = torch.optim.AdamW(model.parameters(), lr = learning_rate)\ncriterion = nn.BCEWithLogitsLoss()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:51.974933Z","iopub.execute_input":"2026-02-26T21:11:51.975305Z","iopub.status.idle":"2026-02-26T21:11:51.979877Z","shell.execute_reply.started":"2026-02-26T21:11:51.975278Z","shell.execute_reply":"2026-02-26T21:11:51.979167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(model)\nx = torch.randn(8, 1, 256, 256).to(device)\ny = model(x)\nprint(y.shape)\nprint(f\"How likely (logits) the model predicts each image in the batch to be pneumonic: \\n{y}\")\n# y.argmax(dim=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T17:55:12.923171Z","iopub.execute_input":"2026-02-26T17:55:12.923820Z","iopub.status.idle":"2026-02-26T17:55:12.954230Z","shell.execute_reply.started":"2026-02-26T17:55:12.923789Z","shell.execute_reply":"2026-02-26T17:55:12.953438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_model_params = 0\nfor param in model.parameters():\n    num_model_params += param.flatten().shape[0]\n\nprint(\"-This Model Has %d (Approximately %d Million) Parameters!\" % (num_model_params, num_model_params//1e6))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T17:55:13.271743Z","iopub.execute_input":"2026-02-26T17:55:13.272039Z","iopub.status.idle":"2026-02-26T17:55:13.278246Z","shell.execute_reply.started":"2026-02-26T17:55:13.272010Z","shell.execute_reply":"2026-02-26T17:55:13.277558Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, roc_curve, auc, confusion_matrix\nimport time","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T17:55:14.890337Z","iopub.execute_input":"2026-02-26T17:55:14.890661Z","iopub.status.idle":"2026-02-26T17:55:14.894567Z","shell.execute_reply.started":"2026-02-26T17:55:14.890632Z","shell.execute_reply":"2026-02-26T17:55:14.893795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, val_loader, device):\n    model.eval()\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in val_loader:\n            images, labels = images.to(device), labels.to(device)\n            labels = labels.float().unsqueeze(1)  # (B,1)\n\n            logits = model(images)\n            probs = torch.sigmoid(logits)\n            preds = (probs > 0.5).long()\n\n            all_preds.append(preds.cpu())\n            all_labels.append(labels.cpu().long())\n\n    all_preds = torch.cat(all_preds)\n    all_labels = torch.cat(all_labels)\n\n    val_acc = accuracy_score(\n        all_labels.numpy(),\n        all_preds.numpy()\n    ) * 100\n\n    return val_acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T13:30:04.567897Z","iopub.execute_input":"2026-01-14T13:30:04.568190Z","iopub.status.idle":"2026-01-14T13:30:04.573577Z","shell.execute_reply.started":"2026-01-14T13:30:04.568162Z","shell.execute_reply":"2026-01-14T13:30:04.572816Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs = 30\n\npatience = 5\ntolerance = 0.005\n\nbest_val_acc = 0.0\nstale_epochs = 0\n\nprint(\"Training started...\")\nfor epoch in range(epochs):\n    model.train()\n    total_loss = 0\n    all_preds = []\n    all_labels = []\n\n    START = time.time()\n    for batch_idx, (images, labels) in enumerate(train_loader):\n        images, labels = images.to(device, non_blocking=True), labels.to(device, non_blocking=True)\n        labels = labels.float().unsqueeze(1)  # (B,1)\n\n        optimizer.zero_grad()\n        logits = model(images)                 # (B,1)\n        \n        loss = criterion(logits, labels)\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n\n        # --- compute predictions ---\n        probs = torch.sigmoid(logits)          # (B,1)\n        preds = (probs > 0.5).long()           # binary 0/1\n        all_preds.append(preds.cpu())\n        all_labels.append(labels.cpu().long())\n\n    TRAIN_TIME = time.time()\n    # --- epoch metrics ---\n    all_preds_tensor = torch.cat(all_preds, dim=0)\n    all_labels_tensor = torch.cat(all_labels, dim=0)\n\n    epoch_acc = accuracy_score(all_labels_tensor.numpy(), all_preds_tensor.numpy()) * 100\n\n    # ---- validation ----\n    val_acc = validate(model, val_loader, device)\n    VAL_TIME = time.time()\n\n    print(\n        f\"===> Epoch {epoch+1}: \"\n        f\"Loss = {total_loss:.4f} | \"\n        f\"Epoch Acc = {epoch_acc:.2f}% | \"\n        f\"Val Acc = {val_acc:.2f}%\"\n    )\n\n    # ---- early stopping ----\n    if val_acc > best_val_acc + tolerance:\n        best_val_acc = val_acc\n        stale_epochs = 0\n\n        torch.save(\n            {\n                \"epoch\": epoch + 1,\n                \"model_state_dict\": model.state_dict(),\n                \"optimizer_state_dict\": optimizer.state_dict(),\n                \"val_acc\": best_val_acc,\n                \"loss\": total_loss,\n            },\n            \"best_model.pth\"\n        )\n    \n        print(f\"Validation improved to {best_val_acc:.2f}%\")\n    else:\n        stale_epochs += 1\n        print(f\"No significant improvement ({stale_epochs}/{patience})\")\n\n    if stale_epochs >= patience:\n        print(\"Early stopping triggered.\")\n        break\n\n    print(f\"Train time: {int((TRAIN_TIME - START)/60.0)}.{int((TRAIN_TIME - START)%60.0)} mins | Val time: {int((VAL_TIME - TRAIN_TIME)/60.0)}.{int((VAL_TIME - TRAIN_TIME)%60.0)} mins\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T13:30:06.311642Z","iopub.execute_input":"2026-01-14T13:30:06.311959Z","iopub.status.idle":"2026-01-14T13:53:10.168488Z","shell.execute_reply.started":"2026-01-14T13:30:06.311925Z","shell.execute_reply":"2026-01-14T13:53:10.164420Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Evaluate**","metadata":{}},{"cell_type":"code","source":"checkpoint = torch.load(\"/kaggle/working/best_model.pth\", map_location=device)\n\nmodel.load_state_dict(checkpoint[\"model_state_dict\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:51.980800Z","iopub.execute_input":"2026-02-26T21:11:51.981078Z","iopub.status.idle":"2026-02-26T21:11:53.659018Z","shell.execute_reply.started":"2026-02-26T21:11:51.981054Z","shell.execute_reply":"2026-02-26T21:11:53.657899Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\nall_preds, all_labels, all_probs = [], [], []\nval_loss = 0.0\n\nwith torch.no_grad():\n    for images, labels in test_loader:\n        images, labels = images.to(device), labels.to(device)\n        labels = labels.float().unsqueeze(1)   # (B,1)\n\n        logits = model(images)                 # (B,1)\n        loss = criterion(logits, labels)\n        val_loss += loss.item()\n\n        probs = torch.sigmoid(logits)          \n        preds = (probs > 0.5).long()           \n\n        all_probs.append(probs.cpu().view(-1))     # (B,)\n        all_preds.append(preds.cpu().view(-1))     # (B,)\n        all_labels.append(labels.cpu().view(-1))   # (B,)\n\n# concatenate\nall_probs = torch.cat(all_probs).numpy()\nall_preds = torch.cat(all_preds).numpy()\nall_labels = torch.cat(all_labels).numpy()\n\n# accuracy\nval_acc = accuracy_score(all_labels, all_preds) * 100\nprint(f\"Validation Loss: {val_loss:.4f} | Validation Accuracy: {val_acc:.2f}%\")\n\n# ROC + AUC\nfpr, tpr, thresholds = roc_curve(all_labels, all_probs)\nroc_auc = auc(fpr, tpr)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.659933Z","iopub.status.idle":"2026-02-26T21:11:53.660315Z","shell.execute_reply.started":"2026-02-26T21:11:53.660169Z","shell.execute_reply":"2026-02-26T21:11:53.660191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_at_threshold(threshold):\n    preds = (all_probs >= threshold).astype(int)\n    acc = accuracy_score(all_labels, preds)\n    return acc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.661636Z","iopub.status.idle":"2026-02-26T21:11:53.661984Z","shell.execute_reply.started":"2026-02-26T21:11:53.661802Z","shell.execute_reply":"2026-02-26T21:11:53.661838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_acc = (0,0)\n\nfor t in [0.3, 0.4, 0.5, 0.6, 0.7]:\n    acc = evaluate_at_threshold(t)\n    if acc > best_acc[0]:\n        best_acc = [acc, t]\n    print(f\"Threshold {t:.1f} → Accuracy: {acc:.4f}\")\n\nthreshold = best_acc[1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.663304Z","iopub.status.idle":"2026-02-26T21:11:53.663586Z","shell.execute_reply.started":"2026-02-26T21:11:53.663446Z","shell.execute_reply":"2026-02-26T21:11:53.663481Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"| Goal                            | TPR (sensitivity) | FPR (1-specificity) |\n| ------------------------------- | ----------------- | ------------------- |\n| Early detection / screening     | High              | Can be higher       |\n| Avoid unnecessary interventions | Moderate          | Low                 |","metadata":{}},{"cell_type":"code","source":"all_preds = (all_probs >= threshold).astype(int)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.664667Z","iopub.status.idle":"2026-02-26T21:11:53.665004Z","shell.execute_reply.started":"2026-02-26T21:11:53.664857Z","shell.execute_reply":"2026-02-26T21:11:53.664883Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import ConfusionMatrixDisplay\n\ncm = confusion_matrix(all_labels, all_preds)\n\ndisp = ConfusionMatrixDisplay(\n    confusion_matrix=cm,\n    display_labels=[\"Negative\", \"Positive\"]\n)\n\ndisp.plot(cmap=\"Blues\", values_format=\"d\")\nplt.title(\"Confusion Matrix\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.665998Z","iopub.status.idle":"2026-02-26T21:11:53.666251Z","shell.execute_reply.started":"2026-02-26T21:11:53.666132Z","shell.execute_reply":"2026-02-26T21:11:53.666150Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"idx_tpr = np.argmax(tpr)\nidx_fpr = np.argmin(fpr)\n\nidx = idx_tpr\n# best_thresh = thresholds[idx]\n\n# print(f\"TPR = {tpr[idx]:.2f} | FPR: {fpr[idx]:.2f}\")\n# print(f\"Best Threshold Value: {best_thresh:.2f}\")\n\nplt.figure(figsize=(6,6))\nplt.plot(fpr, tpr, label=f'ROC curve (AUC = {roc_auc:.3f})')\nplt.plot([0, 1], [0, 1], linestyle='--')\nplt.xlabel('False Positive Rate')\nplt.ylabel('True Positive Rate')\nplt.title('Receiver Operating Characteristic (ROC) Curve')\nplt.legend(loc='lower right')\nplt.grid(True)\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.667305Z","iopub.status.idle":"2026-02-26T21:11:53.667687Z","shell.execute_reply.started":"2026-02-26T21:11:53.667500Z","shell.execute_reply":"2026-02-26T21:11:53.667533Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **Visual Explaination**","metadata":{}},{"cell_type":"code","source":"images, labels = next(iter(test_loader))\nimages, labels = images.to(device), labels.to(device)\n\nprint(images.shape)\nprint(labels.shape)\n\nmodel.eval()\nwith torch.no_grad():\n    logits = model(images)\n    probs = torch.sigmoid(logits)\n    preds = (probs > threshold).long().view(-1)\n\nimages = images.cpu()\nlabels = labels.cpu()\npreds = preds.cpu()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.668425Z","iopub.status.idle":"2026-02-26T21:11:53.668694Z","shell.execute_reply.started":"2026-02-26T21:11:53.668567Z","shell.execute_reply":"2026-02-26T21:11:53.668585Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_images = 16\ncols = 4\nrows = n_images // cols\n\nplt.figure(figsize=(12, 12))\n\nfor i in range(n_images):\n    plt.subplot(rows, cols, i + 1)\n    imshow(images[i])\n    plt.title(f\"Label: {labels[i].item()} | Pred: {preds[i].item()}\")\n    plt.axis(\"off\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.669818Z","iopub.status.idle":"2026-02-26T21:11:53.670257Z","shell.execute_reply.started":"2026-02-26T21:11:53.670020Z","shell.execute_reply":"2026-02-26T21:11:53.670070Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **LIME**","metadata":{}},{"cell_type":"code","source":"!pip install lime -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.671318Z","iopub.status.idle":"2026-02-26T21:11:53.671593Z","shell.execute_reply.started":"2026-02-26T21:11:53.671467Z","shell.execute_reply":"2026-02-26T21:11:53.671485Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from lime import lime_image\n\nexplainer = lime_image.LimeImageExplainer()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.672650Z","iopub.status.idle":"2026-02-26T21:11:53.673025Z","shell.execute_reply.started":"2026-02-26T21:11:53.672872Z","shell.execute_reply":"2026-02-26T21:11:53.672903Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_fn(images): # (N,h,w,3)\n    model.eval()\n    batch_probs = []\n\n    model.eval()\n    images = images[..., 0]  # (N, H, W)\n    images = torch.from_numpy(images).float()\n    images = images.unsqueeze(1)  # (N, 1, H, W)\n\n    images = images.to(device)\n\n    with torch.no_grad():\n        logits = model(images)\n        probs = torch.sigmoid(logits)\n\n    return probs.cpu().numpy()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.674743Z","iopub.status.idle":"2026-02-26T21:11:53.675085Z","shell.execute_reply.started":"2026-02-26T21:11:53.674951Z","shell.execute_reply":"2026-02-26T21:11:53.674973Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from skimage.segmentation import mark_boundaries\n\ndef imshow_lime(img, mask, cmap='gray'):\n        \n    temp_img = mark_boundaries(img, mask)\n    plt.imshow(temp_img, cmap=cmap)\n    plt.axis(\"off\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_images = 9  # how many images to explain\ncols = 3\nrows = n_images // cols\n\nplt.figure(figsize=(12, 8))\n\nfor i in range(n_images):\n   \n    img_np = images[i].numpy().squeeze(0)  # (H, W)\n    \n    explanation = explainer.explain_instance(\n        img_np,\n        predict_fn,\n        top_labels=2,\n        hide_color=0,\n        num_samples=1000\n    );\n    \n    temp, mask = explanation.get_image_and_mask(\n        explanation.top_labels[0],\n        positive_only=False,\n        num_features=5,\n        hide_rest=False\n    )\n    \n    plt.subplot(rows, cols, i + 1)\n    imshow_lime(temp, mask)\n    plt.title(f\"L:{labels[i].item()} | P:{preds[i].item()}\")\n    \nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T21:11:53.677341Z","iopub.status.idle":"2026-02-26T21:11:53.677732Z","shell.execute_reply.started":"2026-02-26T21:11:53.677561Z","shell.execute_reply":"2026-02-26T21:11:53.677590Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **GRAD-CAM**","metadata":{}},{"cell_type":"code","source":"!pip install torchcam --q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T12:11:05.691165Z","iopub.execute_input":"2026-02-26T12:11:05.691497Z","iopub.status.idle":"2026-02-26T12:11:09.049486Z","shell.execute_reply.started":"2026-02-26T12:11:05.691463Z","shell.execute_reply":"2026-02-26T12:11:09.048636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torchcam.methods import SmoothGradCAMpp\nfrom torchcam.utils import overlay_mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-26T12:11:09.051377Z","iopub.execute_input":"2026-02-26T12:11:09.051755Z","iopub.status.idle":"2026-02-26T12:11:09.073071Z","shell.execute_reply.started":"2026-02-26T12:11:09.051719Z","shell.execute_reply":"2026-02-26T12:11:09.072527Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"target_layer = model.layer4[-1]  # last conv layer\ncam_extractor = SmoothGradCAMpp(model, target_layer=target_layer)  # adjust to your model","metadata":{"execution":{"iopub.status.busy":"2026-02-26T12:14:24.287434Z","iopub.execute_input":"2026-02-26T12:14:24.288244Z","iopub.status.idle":"2026-02-26T12:14:24.293301Z","shell.execute_reply.started":"2026-02-26T12:14:24.288209Z","shell.execute_reply":"2026-02-26T12:14:24.292472Z"}}},{"cell_type":"markdown","source":"from PIL import Image\n\ndef show_gradcam(img_tensor, label=None, pred=None):\n    model.eval()\n    logits = model(img_tensor)\n    prob = torch.sigmoid(logits).item()\n    class_idx = 0  # single-output neuron\n\n    # Grad-CAM activation\n    activation_map = cam_extractor(class_idx, logits)[0].squeeze(0).cpu().numpy()\n\n    # Convert to PIL for overlay\n    img_np = img_tensor.squeeze(0).squeeze(0).cpu().numpy()\n    img_np = np.flip(img_np, axis=(0,1))\n    img_pil = Image.fromarray((np.stack([img_np]*3, axis=-1)*255).astype(np.uint8))\n    mask_pil = Image.fromarray((activation_map / activation_map.max() * 255).astype(np.uint8))\n    overlayed = overlay_mask(img_pil, mask_pil, alpha=0.7)\n\n    plt.imshow(overlayed)\n    if label is not None and pred is not None:\n        plt.title(f\"L:{label} | P:{pred}\")\n    plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2026-02-26T12:22:03.710713Z","iopub.execute_input":"2026-02-26T12:22:03.711463Z","iopub.status.idle":"2026-02-26T12:22:03.717910Z","shell.execute_reply.started":"2026-02-26T12:22:03.711428Z","shell.execute_reply":"2026-02-26T12:22:03.717179Z"}}},{"cell_type":"markdown","source":"n_images = 9\ncols = 3\nrows = n_images // cols\n\nplt.figure(figsize=(12, 8))\nfor i in range(n_images):\n    img_tensor = images[i].unsqueeze(0).to(device)\n    plt.subplot(rows, cols, i+1)\n    show_gradcam(img_tensor, label=labels[i].item(), pred=preds[i].item())\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2026-02-26T12:22:04.352986Z","iopub.execute_input":"2026-02-26T12:22:04.353810Z","iopub.status.idle":"2026-02-26T12:22:05.410854Z","shell.execute_reply.started":"2026-02-26T12:22:04.353774Z","shell.execute_reply":"2026-02-26T12:22:05.410098Z"}}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}