{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":22307,"databundleVersionId":1502524,"sourceType":"competition"}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## DICOM JPEG Decompression Support\n\nThe RSNA Pulmonary Embolism dataset contains JPEG-compressed DICOM images.\nTo enable correct decoding of pixel data, we install `pylibjpeg` and\n`pylibjpeg-libjpeg`, which are required by `pydicom` for decompression.\n","metadata":{}},{"cell_type":"code","source":"# Install DICOM JPEG decoders (required for RSNA dataset)\n!pip install -q pylibjpeg pylibjpeg-libjpeg","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:51:07.062715Z","iopub.execute_input":"2026-02-13T11:51:07.063013Z","iopub.status.idle":"2026-02-13T11:51:11.816849Z","shell.execute_reply.started":"2026-02-13T11:51:07.062977Z","shell.execute_reply":"2026-02-13T11:51:11.816025Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport os, cv2\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom glob import glob\nfrom tqdm import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    roc_auc_score, roc_curve,\n    confusion_matrix, accuracy_score,\n    precision_recall_curve, average_precision_score\n)\nimport pydicom\nimport timm\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"Using device:\", DEVICE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:51:31.700982Z","iopub.execute_input":"2026-02-13T11:51:31.70168Z","iopub.status.idle":"2026-02-13T11:51:44.38568Z","shell.execute_reply.started":"2026-02-13T11:51:31.701645Z","shell.execute_reply":"2026-02-13T11:51:44.384887Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_ROOT = \"/kaggle/input/rsna-str-pulmonary-embolism-detection\"\nIMG_ROOT = os.path.join(DATA_ROOT, \"train\")\nLABELS_PATH = os.path.join(DATA_ROOT, \"train.csv\")\n\ndf = pd.read_csv(LABELS_PATH)\n\n# Study-level label\ndf[\"label\"] = (df[\"negative_exam_for_pe\"] == 0).astype(int)\n\nprint(\"Total rows:\", len(df))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:52:11.747967Z","iopub.execute_input":"2026-02-13T11:52:11.748277Z","iopub.status.idle":"2026-02-13T11:52:14.628464Z","shell.execute_reply.started":"2026-02-13T11:52:11.748252Z","shell.execute_reply":"2026-02-13T11:52:14.627844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:52:52.057102Z","iopub.execute_input":"2026-02-13T11:52:52.057639Z","iopub.status.idle":"2026-02-13T11:52:52.080261Z","shell.execute_reply.started":"2026-02-13T11:52:52.057611Z","shell.execute_reply":"2026-02-13T11:52:52.079565Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def series_has_images(study, series):\n    p = os.path.join(IMG_ROOT, study, series)\n    return os.path.exists(p) and len(glob(p+\"/*.dcm\")) > 0\n\nseries_df = df.groupby(\"SeriesInstanceUID\")[\"label\"].first().reset_index()\n\ntrain_ids, val_ids = train_test_split(\n    series_df.SeriesInstanceUID,\n    test_size=0.2,\n    stratify=series_df.label,\n    random_state=42\n)\n\ntrain_df = df[df.SeriesInstanceUID.isin(train_ids)]\nval_df = df[df.SeriesInstanceUID.isin(val_ids)]\n\nprint(\"Train series:\", train_df.SeriesInstanceUID.nunique())\nprint(\"Val series:\", val_df.SeriesInstanceUID.nunique())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:54:21.190533Z","iopub.execute_input":"2026-02-13T11:54:21.190871Z","iopub.status.idle":"2026-02-13T11:54:21.768768Z","shell.execute_reply.started":"2026-02-13T11:54:21.190844Z","shell.execute_reply":"2026-02-13T11:54:21.767999Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sample = train_df.iloc[0]\nsample_path = glob(\n    os.path.join(\n        IMG_ROOT,\n        sample.StudyInstanceUID,\n        sample.SeriesInstanceUID,\n        \"*.dcm\"\n    )\n)[0]\n\ndcm = pydicom.dcmread(sample_path)\nraw = dcm.pixel_array.astype(np.float32)\nraw = raw * dcm.RescaleSlope + dcm.RescaleIntercept\nwin = window_ct(raw)\n\nplt.figure(figsize=(12,4))\nplt.subplot(1,3,1); plt.imshow(raw,cmap=\"gray\"); plt.title(\"Raw CT (HU)\")\nplt.subplot(1,3,2); plt.imshow(win,cmap=\"gray\"); plt.title(\"Windowed CT\")\nplt.subplot(1,3,3); plt.hist(raw.flatten(), bins=200); plt.title(\"HU Histogram\")\nplt.tight_layout(); plt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def window_ct(img, level=100, width=700):\n    low = level - width // 2\n    high = level + width // 2\n    img = np.clip(img, low, high)\n    img = (img - low) / (high - low)\n    return img\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:54:28.632281Z","iopub.execute_input":"2026-02-13T11:54:28.63309Z","iopub.status.idle":"2026-02-13T11:54:28.637311Z","shell.execute_reply.started":"2026-02-13T11:54:28.633055Z","shell.execute_reply":"2026-02-13T11:54:28.636673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self, df, stack=4, train=True):\n        self.groups = df.groupby(\"SeriesInstanceUID\")\n        self.series_ids = list(self.groups.groups.keys())\n        self.stack = stack\n        self.train = train\n\n    def __len__(self):\n        return len(self.series_ids)\n\n    def __getitem__(self, idx):\n        sid = self.series_ids[idx]\n        g = self.groups.get_group(sid)\n\n        study_uid = g.StudyInstanceUID.iloc[0]\n        files = glob(os.path.join(IMG_ROOT, study_uid, sid, \"*.dcm\"))\n\n        slices = []\n        for f in files:\n            dcm = pydicom.dcmread(f)\n            img = dcm.pixel_array.astype(np.float32)\n            img = img * dcm.RescaleSlope + dcm.RescaleIntercept\n            z = float(dcm.ImagePositionPatient[2])\n            img = window_ct(img)\n            img = cv2.resize(img, (224,224))\n            slices.append((z, img))\n\n        slices = [s[1] for s in sorted(slices, key=lambda x: x[0])]\n        n = len(slices)\n\n        if n < 2*self.stack + 1:\n            center = n // 2\n            idxs = [center] * (2*self.stack + 1)\n        else:\n            center = (\n                np.random.randint(self.stack, n-self.stack)\n                if self.train else n//2\n            )\n            idxs = range(center-self.stack, center+self.stack+1)\n\n        x = torch.tensor(np.stack([slices[i] for i in idxs])).unsqueeze(1)\n        y = torch.tensor(g.label.iloc[0], dtype=torch.float32)\n\n        return x, y\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:54:36.186264Z","iopub.execute_input":"2026-02-13T11:54:36.187019Z","iopub.status.idle":"2026-02-13T11:54:36.195149Z","shell.execute_reply.started":"2026-02-13T11:54:36.186988Z","shell.execute_reply":"2026-02-13T11:54:36.194261Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(\n    RSNADataset(train_df, stack=4, train=True),\n    batch_size=4,\n    shuffle=True,\n    num_workers=2\n)\n\nval_loader = DataLoader(\n    RSNADataset(val_df, stack=4, train=False),\n    batch_size=4,\n    shuffle=False,\n    num_workers=2\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:54:46.850766Z","iopub.execute_input":"2026-02-13T11:54:46.851466Z","iopub.status.idle":"2026-02-13T11:54:47.023261Z","shell.execute_reply.started":"2026-02-13T11:54:46.851436Z","shell.execute_reply":"2026-02-13T11:54:47.022503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UnifiedPEModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.encoder = timm.create_model(\n            \"efficientnet_b2\",\n            pretrained=True,\n            in_chans=1,\n            features_only=True\n        )\n\n        C = self.encoder.feature_info[-1][\"num_chs\"]\n\n        self.slice_attn = nn.Sequential(\n            nn.Linear(C, C//2),\n            nn.ReLU(),\n            nn.Linear(C//2, 1)\n        )\n\n        self.cls_head = nn.Sequential(\n            nn.Dropout(0.4),\n            nn.Linear(C, 1)\n        )\n\n        self.seg_head = nn.Sequential(\n            nn.Conv2d(C, C//2, 3, padding=1),\n            nn.ReLU(),\n            nn.Conv2d(C//2, 1, 1)\n        )\n\n        self.feature_maps = None\n        self.feature_grads = {}\n\n    def save_grad(self, idx):\n        def hook(grad):\n            self.feature_grads[idx] = grad\n        return hook\n\n    def forward(self, x):\n\n        B,S,C,H,W = x.shape\n        x = x.view(B*S,C,H,W)\n\n        feats = self.encoder(x)\n\n        self.feature_maps = feats\n        self.feature_grads = {}\n\n        if torch.is_grad_enabled():\n            for i,f in enumerate(feats):\n                f.register_hook(self.save_grad(i))\n\n        feat = feats[-1].view(B,S,feats[-1].shape[1],\n                              feats[-1].shape[2],\n                              feats[-1].shape[3])\n\n        pooled = feat.mean(dim=(3,4))\n        attn = torch.softmax(self.slice_attn(pooled),dim=1)\n        feat = (feat*attn.unsqueeze(-1).unsqueeze(-1)).sum(dim=1)\n\n        cls = self.cls_head(\n            F.adaptive_avg_pool2d(feat,1).flatten(1)\n        ).squeeze(1)\n\n        seg = self.seg_head(feat)\n        seg_up = F.interpolate(seg,(224,224),mode=\"bilinear\")\n\n        return cls, seg_up\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:54:54.311413Z","iopub.execute_input":"2026-02-13T11:54:54.311771Z","iopub.status.idle":"2026-02-13T11:54:54.320385Z","shell.execute_reply.started":"2026-02-13T11:54:54.311746Z","shell.execute_reply":"2026-02-13T11:54:54.319674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiScaleGradCAM:\n    def __init__(self, model, scales=(1,2,3)):\n        self.model = model\n        self.scales = scales\n\n    def generate(self, x):\n        self.model.zero_grad()\n\n        cls, _ = self.model(x)\n        cls.mean().backward(retain_graph=True)\n\n        cams = []\n\n        for i in self.scales:\n            act = self.model.feature_maps[i]\n            grad = self.model.feature_grads[i]\n\n            w = grad.mean(dim=(2,3), keepdim=True)\n            cam = F.relu((w * act).sum(1, keepdim=True))\n            cam = F.interpolate(cam, (224,224))\n            cams.append(cam)\n\n        cam = torch.mean(torch.stack(cams), 0)\n        return cam / (cam.max() + 1e-8)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:56:48.90234Z","iopub.execute_input":"2026-02-13T11:56:48.903142Z","iopub.status.idle":"2026-02-13T11:56:48.908758Z","shell.execute_reply.started":"2026-02-13T11:56:48.90311Z","shell.execute_reply":"2026-02-13T11:56:48.90807Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_pos = (train_df.label == 1).sum()\nnum_neg = (train_df.label == 0).sum()\npos_weight = torch.tensor([num_neg/num_pos]).to(DEVICE)\n\ncls_loss_fn = nn.BCEWithLogitsLoss(pos_weight=pos_weight)\nseg_loss_fn = nn.BCEWithLogitsLoss()\n\ndef dice_loss(pred, target, smooth=1e-6):\n    pred = torch.sigmoid(pred)\n    inter = (pred * target).sum(dim=(1,2,3))\n    union = pred.sum(dim=(1,2,3)) + target.sum(dim=(1,2,3))\n    return 1 - ((2*inter+smooth)/(union+smooth)).mean()\n\nmodel = UnifiedPEModel().to(DEVICE)\ncam_gen = MultiScaleGradCAM(model)\n\noptimizer = torch.optim.AdamW(model.parameters(), lr=3e-5, weight_decay=1e-4)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=50)\nscaler = torch.amp.GradScaler(\"cuda\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:56:54.750452Z","iopub.execute_input":"2026-02-13T11:56:54.751069Z","iopub.status.idle":"2026-02-13T11:56:54.973576Z","shell.execute_reply.started":"2026-02-13T11:56:54.751042Z","shell.execute_reply":"2026-02-13T11:56:54.972847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 50\n\nfor epoch in range(EPOCHS):\n\n    model.train()\n    total = 0\n\n    if epoch < 5:\n        seg_weight = 0.0\n    elif epoch < 15:\n        seg_weight = 0.05\n    else:\n        seg_weight = 0.1\n\n    print(f\"\\nEpoch {epoch+1}/{EPOCHS} | Seg Weight: {seg_weight}\")\n\n    for x,y in tqdm(train_loader):\n        x,y = x.to(DEVICE), y.to(DEVICE)\n        optimizer.zero_grad()\n\n        with torch.amp.autocast(\"cuda\"):\n            cls, seg = model(x)\n            cls_loss = cls_loss_fn(cls,y)\n\n            if seg_weight > 0:\n                cam = cam_gen.generate(x)\n                B,S = x.shape[0],x.shape[1]\n                cam = cam.view(B,S,1,224,224).mean(1)\n\n                cam_min = cam.view(B,-1).min(dim=1)[0].view(B,1,1,1)\n                cam_max = cam.view(B,-1).max(dim=1)[0].view(B,1,1,1)\n                cam_n = (cam-cam_min)/(cam_max-cam_min+1e-8)\n\n                with torch.no_grad():\n                    pseudo_mask = (cam_n>0.35).float()\n\n                pe = y.view(-1,1,1,1)\n                seg_loss = seg_loss_fn(seg*pe,pseudo_mask*pe) + \\\n                           dice_loss(seg*pe,pseudo_mask*pe)\n\n                loss = cls_loss + seg_weight*seg_loss\n            else:\n                loss = cls_loss\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        total += loss.item()\n\n    scheduler.step()\n    print(\"Epoch Loss:\", total/len(train_loader))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-13T11:57:33.68386Z","iopub.execute_input":"2026-02-13T11:57:33.684435Z","iopub.status.idle":"2026-02-13T12:00:28.650878Z","shell.execute_reply.started":"2026-02-13T11:57:33.684406Z","shell.execute_reply":"2026-02-13T12:00:28.649143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\ny_true, y_prob = [], []\n\nwith torch.no_grad():\n    for x, y in tqdm(val_loader):\n        x = x.to(DEVICE)\n        cls, _ = model(x)\n\n        y_true.extend(y.numpy())\n        y_prob.extend(torch.sigmoid(cls).cpu().numpy())\n\ny_true = np.array(y_true)\ny_prob = np.array(y_prob)\n\n# AUROC\nauc = roc_auc_score(y_true, y_prob)\nprint(\"Validation AUROC:\", auc)\n\n# Default threshold 0.5\ny_pred = (y_prob > 0.5).astype(int)\n\ncm = confusion_matrix(y_true, y_pred)\ntn, fp, fn, tp = cm.ravel()\n\nprint(\"Accuracy:\", accuracy_score(y_true, y_pred))\nprint(\"Sensitivity:\", tp/(tp+fn+1e-8))\nprint(\"Specificity:\", tn/(tn+fp+1e-8))\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fpr, tpr, _ = roc_curve(y_true, y_prob)\n\nplt.figure(figsize=(6,5))\nplt.plot(fpr, tpr, label=f\"AUC = {auc:.3f}\")\nplt.plot([0,1], [0,1], '--')\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\nplt.title(\"ROC Curve\")\nplt.legend()\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(5,4))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\")\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"Actual\")\nplt.title(\"Confusion Matrix\")\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"precision, recall, _ = precision_recall_curve(y_true, y_prob)\nap = average_precision_score(y_true, y_prob)\n\nplt.figure(figsize=(6,5))\nplt.plot(recall, precision, label=f\"AP = {ap:.3f}\")\nplt.xlabel(\"Recall (Sensitivity)\")\nplt.ylabel(\"Precision\")\nplt.title(\"Precision–Recall Curve\")\nplt.legend()\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_losses = []\n\nmodel.eval()\nwith torch.no_grad():\n    total = 0\n    for x, y in val_loader:\n        x, y = x.to(DEVICE), y.to(DEVICE)\n        cls, _ = model(x)\n        total += cls_loss_fn(cls, y).item()\n\nval_losses.append(total / len(val_loader))\n\nplt.figure(figsize=(6,5))\nplt.plot(train_losses, label=\"Training Loss\")\nplt.plot(val_losses, label=\"Validation Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training vs Validation Loss\")\nplt.legend()\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"x, y = next(iter(val_loader))\nx = x.to(DEVICE)\n\nmodel.eval()\n\ncls, seg = model(x)\ncam = cam_gen.generate(x)\n\nwith torch.no_grad():\n    seg_out = torch.sigmoid(seg)\n\nB, S = x.shape[0], x.shape[1]\ncam = cam.view(B, S, 1, 224, 224).mean(1)\n\nplt.figure(figsize=(14,4))\n\nplt.subplot(1,3,1)\nplt.imshow(x[0,2,0].cpu(), cmap=\"gray\")\nplt.title(\"CT\")\n\nplt.subplot(1,3,2)\nplt.imshow(cam[0,0].detach().cpu(), cmap=\"jet\")\nplt.title(\"Pseudo Mask (CAM)\")\n\nplt.subplot(1,3,3)\nplt.imshow(seg_out[0,0].cpu(), cmap=\"jet\")\nplt.title(\"Weak Segmentation\")\n\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"thresholds = np.linspace(0, 1, 50)\n\nsens_list, spec_list, acc_list = [], [], []\n\nfor t in thresholds:\n    preds = (y_prob > t).astype(int)\n    tn, fp, fn, tp = confusion_matrix(y_true, preds).ravel()\n\n    sens_list.append(tp/(tp+fn+1e-8))\n    spec_list.append(tn/(tn+fp+1e-8))\n    acc_list.append((tp+tn)/(tp+tn+fp+fn))\n\nplt.figure(figsize=(7,5))\nplt.plot(thresholds, sens_list, label=\"Sensitivity\")\nplt.plot(thresholds, spec_list, label=\"Specificity\")\nplt.plot(thresholds, acc_list, label=\"Accuracy\")\n\nplt.xlabel(\"Decision Threshold\")\nplt.ylabel(\"Metric Value\")\nplt.title(\"Threshold Sensitivity Analysis\")\nplt.legend()\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}