{"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":[{"sourceId":10338,"databundleVersionId":862042,"sourceType":"competition"}],"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-01-14T13:12:49.236484Z","iopub.execute_input":"2026-01-14T13:12:49.236713Z","iopub.status.idle":"2026-01-14T13:12:53.926414Z","shell.execute_reply.started":"2026-01-14T13:12:49.236691Z","shell.execute_reply":"2026-01-14T13:12:53.925417Z"}},"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-01-14T13:12:53.928439Z","iopub.execute_input":"2026-01-14T13:12:53.928738Z","iopub.status.idle":"2026-01-14T13:13:08.309829Z","shell.execute_reply.started":"2026-01-14T13:12:53.928702Z","shell.execute_reply":"2026-01-14T13:13:08.309212Z"}},"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-01-14T13:13:08.313091Z","iopub.execute_input":"2026-01-14T13:13:08.313284Z","iopub.status.idle":"2026-01-14T13:13:08.409376Z","shell.execute_reply.started":"2026-01-14T13:13:08.313262Z","shell.execute_reply":"2026-01-14T13:13:08.408647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"labels_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T13:21:46.822299Z","iopub.execute_input":"2026-01-14T13:21:46.822921Z","iopub.status.idle":"2026-01-14T13:21:46.84018Z","shell.execute_reply.started":"2026-01-14T13:21:46.822884Z","shell.execute_reply":"2026-01-14T13:21:46.839644Z"}},"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-01-14T13:21:47.349785Z","iopub.execute_input":"2026-01-14T13:21:47.350066Z","iopub.status.idle":"2026-01-14T13:21:48.170532Z","shell.execute_reply.started":"2026-01-14T13:21:47.35004Z","shell.execute_reply":"2026-01-14T13:21:48.169786Z"}},"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-01-14T13:21:48.678Z","iopub.execute_input":"2026-01-14T13:21:48.678293Z","iopub.status.idle":"2026-01-14T13:21:48.726941Z","shell.execute_reply.started":"2026-01-14T13:21:48.678265Z","shell.execute_reply":"2026-01-14T13:21:48.726368Z"}},"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-01-14T13:21:54.05999Z","iopub.execute_input":"2026-01-14T13:21:54.060284Z","iopub.status.idle":"2026-01-14T13:21:54.064674Z","shell.execute_reply.started":"2026-01-14T13:21:54.060253Z","shell.execute_reply":"2026-01-14T13:21:54.063875Z"}},"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-01-14T13:22:00.790045Z","iopub.execute_input":"2026-01-14T13:22:00.790567Z","iopub.status.idle":"2026-01-14T13:28:11.020327Z","shell.execute_reply.started":"2026-01-14T13:22:00.790537Z","shell.execute_reply":"2026-01-14T13:28:11.019613Z"}},"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-01-14T13:28:11.022017Z","iopub.execute_input":"2026-01-14T13:28:11.022312Z","iopub.status.idle":"2026-01-14T13:28:11.057491Z","shell.execute_reply.started":"2026-01-14T13:28:11.022286Z","shell.execute_reply":"2026-01-14T13:28:11.056889Z"}},"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-01-14T13:28:11.058356Z","iopub.execute_input":"2026-01-14T13:28:11.058645Z","iopub.status.idle":"2026-01-14T13:28:11.466857Z","shell.execute_reply.started":"2026-01-14T13:28:11.058621Z","shell.execute_reply":"2026-01-14T13:28:11.466244Z"}},"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-01-14T13:28:11.468384Z","iopub.execute_input":"2026-01-14T13:28:11.468667Z","iopub.status.idle":"2026-01-14T13:28:11.476101Z","shell.execute_reply.started":"2026-01-14T13:28:11.468641Z","shell.execute_reply":"2026-01-14T13:28:11.47548Z"}},"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-01-14T13:28:11.476854Z","iopub.execute_input":"2026-01-14T13:28:11.477136Z","iopub.status.idle":"2026-01-14T13:28:11.483071Z","shell.execute_reply.started":"2026-01-14T13:28:11.477087Z","shell.execute_reply":"2026-01-14T13:28:11.482352Z"}},"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-01-14T13:28:50.441952Z","iopub.execute_input":"2026-01-14T13:28:50.442254Z","iopub.status.idle":"2026-01-14T13:28:50.550985Z","shell.execute_reply.started":"2026-01-14T13:28:50.442224Z","shell.execute_reply":"2026-01-14T13:28:50.55026Z"}},"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-01-14T13:28:51.543595Z","iopub.execute_input":"2026-01-14T13:28:51.544235Z","iopub.status.idle":"2026-01-14T13:28:51.55209Z","shell.execute_reply.started":"2026-01-14T13:28:51.544197Z","shell.execute_reply":"2026-01-14T13:28:51.551461Z"}},"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) for s in batch])\n    labels = torch.tensor([s.Label.item() for s in batch])\n    return images, labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T13:28:52.823816Z","iopub.execute_input":"2026-01-14T13:28:52.824112Z","iopub.status.idle":"2026-01-14T13:28:52.828697Z","shell.execute_reply.started":"2026-01-14T13:28:52.824082Z","shell.execute_reply":"2026-01-14T13:28:52.828039Z"}},"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-01-14T13:28:53.929053Z","iopub.execute_input":"2026-01-14T13:28:53.930079Z","iopub.status.idle":"2026-01-14T13:28:53.934229Z","shell.execute_reply.started":"2026-01-14T13:28:53.930033Z","shell.execute_reply":"2026-01-14T13:28:53.933548Z"}},"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-01-14T13:28:55.63995Z","iopub.execute_input":"2026-01-14T13:28:55.64049Z","iopub.status.idle":"2026-01-14T13:29:15.559112Z","shell.execute_reply.started":"2026-01-14T13:28:55.640457Z","shell.execute_reply":"2026-01-14T13:29:15.558146Z"}},"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-01-14T13:29:15.561612Z","iopub.execute_input":"2026-01-14T13:29:15.562185Z","iopub.status.idle":"2026-01-14T13:29:15.568328Z","shell.execute_reply.started":"2026-01-14T13:29:15.562135Z","shell.execute_reply":"2026-01-14T13:29:15.567525Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision.models as models\n\nmodel = models.regnet_y_400mf(pretrained=True)\n\nmodel.stem[0] = nn.Conv2d(1, model.stem[0].out_channels, kernel_size=3, stride=2, padding=1, bias=False)\n\nmodel.fc = nn.Linear(model.fc.in_features, 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T13:29:23.832229Z","iopub.execute_input":"2026-01-14T13:29:23.832786Z","iopub.status.idle":"2026-01-14T13:29:24.399529Z","shell.execute_reply.started":"2026-01-14T13:29:23.832745Z","shell.execute_reply":"2026-01-14T13:29:24.398952Z"}},"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-01-14T13:29:46.681699Z","iopub.execute_input":"2026-01-14T13:29:46.682426Z","iopub.status.idle":"2026-01-14T13:29:46.687247Z","shell.execute_reply.started":"2026-01-14T13:29:46.682379Z","shell.execute_reply":"2026-01-14T13:29:46.686515Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.to(device);","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T13:29:48.103696Z","iopub.execute_input":"2026-01-14T13:29:48.104484Z","iopub.status.idle":"2026-01-14T13:29:48.133043Z","shell.execute_reply.started":"2026-01-14T13:29:48.104418Z","shell.execute_reply":"2026-01-14T13:29:48.132353Z"}},"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-01-14T13:29:50.039521Z","iopub.execute_input":"2026-01-14T13:29:50.040101Z","iopub.status.idle":"2026-01-14T13:29:50.044637Z","shell.execute_reply.started":"2026-01-14T13:29:50.040066Z","shell.execute_reply":"2026-01-14T13:29:50.043962Z"}},"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-01-14T13:29:52.183098Z","iopub.execute_input":"2026-01-14T13:29:52.183395Z","iopub.status.idle":"2026-01-14T13:29:53.053625Z","shell.execute_reply.started":"2026-01-14T13:29:52.183366Z","shell.execute_reply":"2026-01-14T13:29:53.052887Z"}},"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-01-14T13:30:03.128071Z","iopub.execute_input":"2026-01-14T13:30:03.128373Z","iopub.status.idle":"2026-01-14T13:30:03.137807Z","shell.execute_reply.started":"2026-01-14T13:30:03.128344Z","shell.execute_reply":"2026-01-14T13:30:03.137109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import accuracy_score, roc_curve, auc\nimport time","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-14T13:30:03.443903Z","iopub.execute_input":"2026-01-14T13:30:03.444208Z","iopub.status.idle":"2026-01-14T13:30:03.448215Z","shell.execute_reply.started":"2026-01-14T13:30:03.444177Z","shell.execute_reply":"2026-01-14T13:30:03.447507Z"}},"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.56819Z","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.16442Z"}},"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},"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-01-14T13:53:17.018975Z","iopub.execute_input":"2026-01-14T13:53:17.020031Z","iopub.status.idle":"2026-01-14T13:54:44.254798Z","shell.execute_reply.started":"2026-01-14T13:53:17.019978Z","shell.execute_reply":"2026-01-14T13:54:44.253869Z"}},"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":"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-01-14T13:54:44.256765Z","iopub.execute_input":"2026-01-14T13:54:44.257124Z","iopub.status.idle":"2026-01-14T13:54:44.396263Z","shell.execute_reply.started":"2026-01-14T13:54:44.25709Z","shell.execute_reply":"2026-01-14T13:54:44.395693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}