{"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-09T10:01:03.710359Z","iopub.execute_input":"2026-02-09T10:01:03.710641Z","iopub.status.idle":"2026-02-09T10:01:08.644748Z","shell.execute_reply.started":"2026-02-09T10:01:03.710614Z","shell.execute_reply":"2026-02-09T10:01:08.643872Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\n\nprint(\"CUDA available:\", torch.cuda.is_available())\nprint(\"CUDA device count:\", torch.cuda.device_count())\n\nif torch.cuda.is_available():\n    print(\"Current GPU:\", torch.cuda.get_device_name(0))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:11.98197Z","iopub.execute_input":"2026-02-09T10:01:11.982281Z","iopub.status.idle":"2026-02-09T10:01:16.287672Z","shell.execute_reply.started":"2026-02-09T10:01:11.982244Z","shell.execute_reply":"2026-02-09T10:01:16.287039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os, cv2, torch, timm, pydicom\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\n\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 roc_auc_score, roc_curve, confusion_matrix\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:17.661462Z","iopub.execute_input":"2026-02-09T10:01:17.662124Z","iopub.status.idle":"2026-02-09T10:01:27.187947Z","shell.execute_reply.started":"2026-02-09T10:01:17.662094Z","shell.execute_reply":"2026-02-09T10:01:27.187216Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"Device:\", DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:29.862221Z","iopub.execute_input":"2026-02-09T10:01:29.86298Z","iopub.status.idle":"2026-02-09T10:01:29.867084Z","shell.execute_reply.started":"2026-02-09T10:01:29.862951Z","shell.execute_reply":"2026-02-09T10:01:29.866506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Kaggle input paths\nDATA_ROOT = \"/kaggle/input/rsna-str-pulmonary-embolism-detection\"\nIMG_ROOT = os.path.join(DATA_ROOT, \"train\")\nCSV_PATH = os.path.join(DATA_ROOT, \"train.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:31.398948Z","iopub.execute_input":"2026-02-09T10:01:31.399259Z","iopub.status.idle":"2026-02-09T10:01:31.403634Z","shell.execute_reply.started":"2026-02-09T10:01:31.399233Z","shell.execute_reply":"2026-02-09T10:01:31.402722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(CSV_PATH)\n\n# Exam-level PE label\ndf[\"label\"] = (df[\"negative_exam_for_pe\"] == 0).astype(int)\n\nprint(\"Total rows:\", len(df))\ndf.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:34.03Z","iopub.execute_input":"2026-02-09T10:01:34.030496Z","iopub.status.idle":"2026-02-09T10:01:36.925219Z","shell.execute_reply.started":"2026-02-09T10:01:34.030471Z","shell.execute_reply":"2026-02-09T10:01:36.924636Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = df[df.SeriesInstanceUID.notnull()]\nprint(\"Unique series:\", df.SeriesInstanceUID.nunique())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:38.731262Z","iopub.execute_input":"2026-02-09T10:01:38.731996Z","iopub.status.idle":"2026-02-09T10:01:39.186854Z","shell.execute_reply.started":"2026-02-09T10:01:38.731967Z","shell.execute_reply":"2026-02-09T10:01:39.186177Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:40.335848Z","iopub.execute_input":"2026-02-09T10:01:40.336385Z","iopub.status.idle":"2026-02-09T10:01:40.556496Z","shell.execute_reply.started":"2026-02-09T10:01:40.336359Z","shell.execute_reply":"2026-02-09T10:01:40.555697Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_df = df.groupby(\"SeriesInstanceUID\")[\"label\"].first().reset_index()\n\nsns.countplot(x=\"label\", data=series_df)\nplt.title(\"PE vs Non-PE Distribution (Subset)\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:42.065187Z","iopub.execute_input":"2026-02-09T10:01:42.066078Z","iopub.status.idle":"2026-02-09T10:01:42.382631Z","shell.execute_reply.started":"2026-02-09T10:01:42.066049Z","shell.execute_reply":"2026-02-09T10:01:42.382152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"series_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)]\nprint(\"Train series:\", train_df.SeriesInstanceUID.nunique())\nprint(\"Val series:\", val_df.SeriesInstanceUID.nunique())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:44.726012Z","iopub.execute_input":"2026-02-09T10:01:44.726334Z","iopub.status.idle":"2026-02-09T10:01:45.285181Z","shell.execute_reply.started":"2026-02-09T10:01:44.726294Z","shell.execute_reply":"2026-02-09T10:01:45.284376Z"}},"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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:47.084102Z","iopub.execute_input":"2026-02-09T10:01:47.084444Z","iopub.status.idle":"2026-02-09T10:01:47.088265Z","shell.execute_reply.started":"2026-02-09T10:01:47.084419Z","shell.execute_reply":"2026-02-09T10:01:47.087774Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:01:49.724087Z","iopub.execute_input":"2026-02-09T10:01:49.724435Z","iopub.status.idle":"2026-02-09T10:01:49.727924Z","shell.execute_reply.started":"2026-02-09T10:01:49.724407Z","shell.execute_reply":"2026-02-09T10:01:49.727333Z"}},"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,"execution":{"iopub.status.busy":"2026-02-09T10:01:52.364395Z","iopub.execute_input":"2026-02-09T10:01:52.364681Z","iopub.status.idle":"2026-02-09T10:01:53.04509Z","shell.execute_reply.started":"2026-02-09T10:01:52.364657Z","shell.execute_reply":"2026-02-09T10:01:53.044339Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self, df, stack=2, 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        study_uid = g.StudyInstanceUID.iloc[0]\n\n        files = glob(os.path.join(IMG_ROOT, study_uid, sid, \"*.dcm\"))\n        slices = []\n\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-09T10:01:55.94467Z","iopub.execute_input":"2026-02-09T10:01:55.945496Z","iopub.status.idle":"2026-02-09T10:01:55.955165Z","shell.execute_reply.started":"2026-02-09T10:01:55.945463Z","shell.execute_reply":"2026-02-09T10:01:55.954382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(\n    RSNADataset(train_df, train=True),\n    batch_size=4,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    RSNADataset(val_df, train=False),\n    batch_size=4,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:02:02.768695Z","iopub.execute_input":"2026-02-09T10:02:02.769286Z","iopub.status.idle":"2026-02-09T10:02:03.160708Z","shell.execute_reply.started":"2026-02-09T10:02:02.76926Z","shell.execute_reply":"2026-02-09T10:02:03.160086Z"}},"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_b0\",\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.Linear(C,1)\n        self.det_head = nn.Conv2d(C,1,1)\n\n        self.feature_maps = None\n        self.feature_grads = {}\n\n    def save_grad(self, idx):\n        def hook(grad): self.feature_grads[idx] = grad\n        return hook\n\n    def forward(self, x):\n        B,S,C,H,W = x.shape\n        x = x.view(B*S,C,H,W)\n\n        feats = self.encoder(x)\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(\n            B,S,feats[-1].shape[1],\n            feats[-1].shape[2],\n            feats[-1].shape[3]\n        )\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        det = self.det_head(feat)\n        return cls, det\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:02:05.626412Z","iopub.execute_input":"2026-02-09T10:02:05.627018Z","iopub.status.idle":"2026-02-09T10:02:05.634823Z","shell.execute_reply.started":"2026-02-09T10:02:05.626991Z","shell.execute_reply":"2026-02-09T10:02:05.634086Z"}},"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        cls,_ = self.model(x)\n        cls.mean().backward(retain_graph=True)\n\n        cams=[]\n        for i in self.scales:\n            act = self.model.feature_maps[i]\n            grad = self.model.feature_grads[i]\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-09T10:02:09.680036Z","iopub.execute_input":"2026-02-09T10:02:09.680645Z","iopub.status.idle":"2026-02-09T10:02:09.686206Z","shell.execute_reply.started":"2026-02-09T10:02:09.680616Z","shell.execute_reply":"2026-02-09T10:02:09.685415Z"}},"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)\ndet_loss_fn = nn.BCEWithLogitsLoss()\nmse_loss_fn = nn.MSELoss()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:02:13.001507Z","iopub.execute_input":"2026-02-09T10:02:13.002041Z","iopub.status.idle":"2026-02-09T10:02:13.254606Z","shell.execute_reply.started":"2026-02-09T10:02:13.002011Z","shell.execute_reply":"2026-02-09T10:02:13.254044Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UnifiedPEModel().to(DEVICE)\ncam_gen = MultiScaleGradCAM(model)\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=20\n)\n\n# ✅ Updated AMP API (no deprecation warning)\nscaler = torch.amp.GradScaler(\"cuda\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-09T10:02:15.912048Z","iopub.execute_input":"2026-02-09T10:02:15.91277Z","iopub.status.idle":"2026-02-09T10:02:16.850981Z","shell.execute_reply.started":"2026-02-09T10:02:15.912739Z","shell.execute_reply":"2026-02-09T10:02:16.850454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 20\ntrain_losses=[]\n\nfor epoch in range(EPOCHS):\n    model.train()\n    total=0\n    lambda_cons = 0.0 if epoch < 3 else 0.3\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.cuda.amp.autocast():\n            cls, det = model(x)\n\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            det_up = F.interpolate(det,(224,224))\n            cam_n = (cam-cam.min())/(cam.max()-cam.min()+1e-8)\n            det_n = (det_up-det_up.min())/(det_up.max()-det_up.min()+1e-8)\n\n            pe = y.view(-1,1,1,1)\n            loss = (\n                cls_loss_fn(cls,y) +\n                det_loss_fn(det_up*pe,(cam_n>0.35).float()*pe) +\n                lambda_cons*mse_loss_fn(cam_n*pe, det_n*pe)\n            )\n\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        total += loss.item()\n\n    scheduler.step()\n    train_losses.append(total/len(train_loader))\n    print(f\"Epoch {epoch+1}/{EPOCHS} | Loss {train_losses[-1]:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}