{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"jupytext":{"cell_metadata_filter":"-all","main_language":"python","notebook_metadata_filter":"-all"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13851420,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13377158,"sourceType":"datasetVersion","datasetId":8486993},{"sourceId":13377243,"sourceType":"datasetVersion","datasetId":8487053},{"sourceId":13377355,"sourceType":"datasetVersion","datasetId":8487114},{"sourceId":13377471,"sourceType":"datasetVersion","datasetId":8487193},{"sourceId":13377583,"sourceType":"datasetVersion","datasetId":8487279},{"sourceId":13377743,"sourceType":"datasetVersion","datasetId":8487350},{"sourceId":13377903,"sourceType":"datasetVersion","datasetId":8487483},{"sourceId":13378062,"sourceType":"datasetVersion","datasetId":8487615},{"sourceId":13378221,"sourceType":"datasetVersion","datasetId":8487733},{"sourceId":13378367,"sourceType":"datasetVersion","datasetId":8487845},{"sourceId":13378521,"sourceType":"datasetVersion","datasetId":8487955},{"sourceId":13378653,"sourceType":"datasetVersion","datasetId":8488077},{"sourceId":13378772,"sourceType":"datasetVersion","datasetId":8488187},{"sourceId":13378916,"sourceType":"datasetVersion","datasetId":8488267},{"sourceId":13379037,"sourceType":"datasetVersion","datasetId":8488367},{"sourceId":13379234,"sourceType":"datasetVersion","datasetId":8488460},{"sourceId":13380531,"sourceType":"datasetVersion","datasetId":8489585}],"dockerImageVersionId":31089,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Cell 1\nimport os\nimport glob\nimport random\nimport shutil\nimport warnings\nfrom collections import defaultdict\n\nimport pandas as pd\nimport polars as pl\nimport pydicom\n\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision.models import resnet18, ResNet18_Weights\nimport torchvision.transforms.functional as TF\nfrom torchvision.transforms import InterpolationMode  # enum for interpolation\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score\n\nimport kaggle_evaluation.rsna_inference_server\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# keep libs single-threaded for stability\nos.environ[\"OMP_NUM_THREADS\"] = \"1\"\ncv2.setNumThreads(0)\ntorch.set_num_threads(1)\n\n# CUDNN determinism (submission-friendly)\ntorch.backends.cudnn.benchmark = False\ntorch.backends.cudnn.deterministic = True\n\n# Debug verbosity (set to \"0\" to reduce logs)\nDEBUG = int(os.environ.get(\"RSNA_DEBUG\", \"1\"))\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)\n\ndef log(msg: str):\n    if DEBUG:\n        print(f\"[RSNA] {msg}\", flush=True)\n\ndef log_once(label, cache=set()):\n    def _inner(msg):\n        key = f\"{label}:{msg}\"\n        if key not in cache:\n            cache.add(key)\n            log(msg)\n    return _inner\n\nlog_warn = log_once(\"warn\")\nlog_err  = log_once(\"err\")\n\n# quieter pydicom warnings\nwarnings.filterwarnings(\"ignore\", message=\".*invalid value for VR UI.*\", category=UserWarning)\nwarnings.filterwarnings(\"ignore\", message=\".*Number of Frames.*\", category=UserWarning)\n\nlog(f\"Device: {device.type} | CUDA: {torch.cuda.is_available()}\")\nif torch.cuda.is_available():\n    props = torch.cuda.get_device_properties(0)\n    log(f\"GPU: {props.name}, {props.total_memory/1024**3:.1f} GB total\")\n\nos.environ[\"TORCH_HOME\"] = \"/kaggle/working/torch_home\"  # never touch ~/.cache\ntorch_hub_dir = os.path.join(os.environ[\"TORCH_HOME\"], \"hub\")\ntorch_hub_ckpt = os.path.join(torch_hub_dir, \"checkpoints\", \"resnet18-f37072fd.pth\")\n\nsrc_candidates = [\n    \"/kaggle/input/weights/checkpoints/resnet18-f37072fd.pth\",\n    \"/kaggle/input/weights/resnet18-f37072fd.pth\",\n]\n\nos.makedirs(os.path.dirname(torch_hub_ckpt), exist_ok=True)\nfor src in src_candidates:\n    if os.path.exists(src):\n        try:\n            shutil.copy(src, torch_hub_ckpt)\n            break\n        except Exception as e:\n            log_warn(f\"Copying pretrained weight failed from {src}: {e}\")\n\ntry:\n    import torch.hub as th\n    th.set_dir(torch_hub_dir)\nexcept Exception as e:\n    log_warn(f\"torch.hub.set_dir failed: {e}\")\n\nif os.path.exists(torch_hub_ckpt):\n    log(f\"ImageNet weights cached at {torch_hub_ckpt}\")\nelse:\n    log_warn(\"resnet18-f37072fd.pth not found in /kaggle/input/weights — will fall back to random init if needed.\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T10:01:30.083572Z","iopub.execute_input":"2025-10-14T10:01:30.083927Z","iopub.status.idle":"2025-10-14T10:01:45.854661Z","shell.execute_reply.started":"2025-10-14T10:01:30.083908Z","shell.execute_reply":"2025-10-14T10:01:45.854061Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 2\nID_COL = 'SeriesInstanceUID'\nLABEL_COLS = [\n    'Left Infraclinoid Internal Carotid Artery',\n    'Right Infraclinoid Internal Carotid Artery',\n    'Left Supraclinoid Internal Carotid Artery',\n    'Right Supraclinoid Internal Carotid Artery',\n    'Left Middle Cerebral Artery',\n    'Right Middle Cerebral Artery',\n    'Anterior Communicating Artery',\n    'Left Anterior Cerebral Artery',\n    'Right Anterior Cerebral Artery',\n    'Left Posterior Communicating Artery',\n    'Right Posterior Communicating Artery',\n    'Basilar Tip',\n    'Other Posterior Circulation',\n    'Aneurysm Present',\n]\n\nDICOM_TAG_ALLOWLIST = [\n    'BitsAllocated','BitsStored','Columns','FrameOfReferenceUID','HighBit',\n    'ImageOrientationPatient','ImagePositionPatient','InstanceNumber','Modality',\n    'PatientID','PhotometricInterpretation','PixelRepresentation','PixelSpacing',\n    'PlanarConfiguration','RescaleIntercept','RescaleSlope','RescaleType','Rows',\n    'SOPClassUID','SOPInstanceUID','SamplesPerPixel','SliceThickness',\n    'SpacingBetweenSlices','StudyInstanceUID','TransferSyntaxUID',\n]\n\nDATA_ROOT   = \"/kaggle/input/rsna-intracranial-aneurysm-detection\"\nSERIES_ROOT = f\"{DATA_ROOT}/series\"\nTRAIN_CSV   = f\"{DATA_ROOT}/train.csv\"\n\nPREPROCESSED_ROOTS = []\ntry:\n    for p in os.listdir(\"/kaggle/input\"):\n        plow = p.lower()\n        if (\n            plow.startswith(\"rsna-preprocessed-v2-shard\") or\n            plow.startswith(\"rsna-preprocessed-shards-v2\")  # in case of alternate naming\n        ):\n            PREPROCESSED_ROOTS.append(os.path.join(\"/kaggle/input\", p))\nexcept Exception as e:\n    log_warn(f\"Listing /kaggle/input failed: {e}\")\n\nlog(f\"Discovered {len(PREPROCESSED_ROOTS)} preprocessed roots:\")\nfor r in PREPROCESSED_ROOTS:\n    log(f\"  - {r}\")\n\n_initialized = False\nmodel = None\n\nIS_SUBMISSION = bool(os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"))\nlog(f\"IS_SUBMISSION={IS_SUBMISSION}\")\n\nSUBSET_N  = 4348 if not IS_SUBMISSION else 4000\nEPOCHS    = 9     if not IS_SUBMISSION else 6\nBATCH_SZ  = 8\nK_SLICES  = 9\nIMG_SIZE  = 224\n\nCOL_WEIGHTS = torch.ones(14, dtype=torch.float32)\nCOL_WEIGHTS[-1] = 4.0\n\nAP_POS_WEIGHT = None\n\nAUG_PROB_HFLIP = 0.5\nAUG_PROB_VFLIP = 0.2\nAUG_BRIGHTNESS = 0.10\nAUG_CONTRAST   = 0.10\n\nWINDOW_MIN, WINDOW_MAX = -100.0, 300.0\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T10:01:45.855336Z","iopub.execute_input":"2025-10-14T10:01:45.855678Z","iopub.status.idle":"2025-10-14T10:01:45.881298Z","shell.execute_reply.started":"2025-10-14T10:01:45.855659Z","shell.execute_reply":"2025-10-14T10:01:45.880575Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 3\ndef preprocess_dicom(filepath, img_size=IMG_SIZE):\n    try:\n        dcm = pydicom.dcmread(filepath)\n        img = dcm.pixel_array\n    except Exception:\n        dcm = pydicom.dcmread(filepath, force=True)\n        try:\n            img = dcm.pixel_array\n        except Exception:\n            img = np.zeros((512, 512), dtype=np.float32)\n\n    img = np.asarray(img)\n    if img.size == 0:\n        img = np.zeros((512, 512), dtype=np.float32)\n\n    if img.ndim == 3:\n        if img.shape[-1] in (3,4):  # H,W,C\n            if img.shape[-1] == 4: img = img[..., :3]\n            img = cv2.cvtColor(img.astype(np.float32), cv2.COLOR_RGB2GRAY)\n        else:  # F,H,W -> middle frame\n            img = img[img.shape[0] // 2]\n\n    img = np.squeeze(img).astype(np.float32)\n    if img.ndim != 2:\n        img = np.zeros((512, 512), dtype=np.float32)\n\n    img = np.clip(img, WINDOW_MIN, WINDOW_MAX)\n    mn, mx = float(img.min()), float(img.max())\n    rng = mx - mn\n    img = np.zeros_like(img) if rng < 1e-5 else (img - mn)/(rng + 1e-5)\n\n    img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_AREA)\n    img_t = torch.from_numpy(img).float().unsqueeze(0).expand(3, -1, -1).contiguous()\n    return img_t  # [3,H,W]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T10:01:45.882829Z","iopub.execute_input":"2025-10-14T10:01:45.883102Z","iopub.status.idle":"2025-10-14T10:01:45.907612Z","shell.execute_reply.started":"2025-10-14T10:01:45.883085Z","shell.execute_reply":"2025-10-14T10:01:45.906952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 4\ndef pick_k_indices(n, k):\n    if n <= 0:\n        log_warn(\"pick_k_indices called with n<=0; returning zeros\")\n        return np.zeros(k, dtype=int)\n    idxs = np.linspace(0, n - 1, num=k)\n    idxs = np.round(idxs).astype(int)\n    idxs = np.clip(idxs, 0, n - 1)\n    return idxs\n\ndef _apply_tensor_aug(xk: torch.Tensor) -> torch.Tensor:\n    if torch.rand(1).item() < AUG_PROB_HFLIP:\n        xk = torch.flip(xk, dims=[3])\n    if torch.rand(1).item() < AUG_PROB_VFLIP:\n        xk = torch.flip(xk, dims=[2])\n    if AUG_BRIGHTNESS > 0 or AUG_CONTRAST > 0:\n        xf = xk.float()\n        if AUG_BRIGHTNESS > 0:\n            delta = (torch.rand(1).item() * 2 - 1) * AUG_BRIGHTNESS\n            xf = xf + delta\n        if AUG_CONTRAST > 0:\n            alpha = 1.0 + ((torch.rand(1).item() * 2 - 1) * AUG_CONTRAST)\n            xf = (xf - 0.5) * alpha + 0.5\n        xk = xf.clamp(0,1).to(xk.dtype)\n    return xk\n\nclass RSNATensorDataset(Dataset):\n    def __init__(self, df, series_root, preproc_roots, k_slices=K_SLICES, img_size=IMG_SIZE, train=True):\n        self.df = df.reset_index(drop=True)\n        self.series_root = series_root\n        self.preproc_roots = preproc_roots\n        self.k = k_slices\n        self.img_size = img_size\n        self.train = train\n\n        def _sid_from_pt(pt):\n            return os.path.basename(pt)[:-3]\n\n        sid2path = {}\n        flat_cnt = 0\n        shard_cnt = 0\n\n        for root in self.preproc_roots:\n            flat_k = os.path.join(root, \"pre_k\")\n            if os.path.isdir(flat_k):\n                pts = glob.glob(os.path.join(flat_k, \"*.pt\"))\n                for kpt in pts:\n                    sid2path[_sid_from_pt(kpt)] = kpt\n                flat_cnt += len(pts)\n                log(f\"Indexed {len(pts)} pre_k PT files in flat dir: {flat_k}\")\n\n            # optional nested layout\n            for sd in [d for d in os.listdir(root) if d.lower().startswith(\"shard\")]:\n                kdir = os.path.join(root, sd, \"pre_k\")\n                if os.path.isdir(kdir):\n                    pts = glob.glob(os.path.join(kdir, \"*.pt\"))\n                    for kpt in pts:\n                        sid2path[_sid_from_pt(kpt)] = kpt\n                    shard_cnt += len(pts)\n                    log(f\"Indexed {len(pts)} pre_k PT files in shard dir: {kdir}\")\n\n        self.sid2path = sid2path\n        log(f\"Total indexed pre_k tensors: {len(self.sid2path)} (flat={flat_cnt}, sharded={shard_cnt})\")\n\n        sids = set(self.df['SeriesInstanceUID'].astype(str).tolist())\n        hits = sum(1 for s in sids if s in self.sid2path)\n        miss = len(sids) - hits\n        log(f\"Dataset SIDs: {len(sids)} | with PT={hits} | fallback raw DICOM={miss}\")\n\n    def __len__(self): \n        return len(self.df)\n\n    def _load_pre_k(self, sid: str):\n        p = self.sid2path.get(sid, None)\n        if p and os.path.isfile(p):\n            try:\n                t = torch.load(p, map_location=\"cpu\")\n                if t.dtype != torch.float32:\n                    t = t.float()\n                return torch.clamp(t, 0, 1)\n            except Exception as e:\n                log_warn(f\"torch.load failed for {os.path.basename(p)}: {e}\")\n                return None\n        return None\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        sid = str(row['SeriesInstanceUID'])\n        y   = torch.tensor(row[LABEL_COLS].values.astype(np.float32))\n\n        X = self._load_pre_k(sid)\n        if X is None:\n            sp = os.path.join(self.series_root, sid)\n            try:\n                files = sorted([f for f in os.listdir(sp) if f.lower().endswith(\".dcm\")])\n            except Exception as e:\n                log_err(f\"listdir failed for {sp}: {e}\")\n                files = []\n            if len(files) == 0:\n                log_err(f\"No DICOMs for {sid} — returning zeros\")\n                X = torch.zeros(self.k, 3, self.img_size, self.img_size, dtype=torch.float32)\n            else:\n                idxs = pick_k_indices(len(files), self.k)\n                paths = [os.path.join(sp, files[i]) for i in idxs]\n                imgs  = [preprocess_dicom(p, self.img_size) for p in paths]\n                X     = torch.stack(imgs, dim=0).float().clamp(0,1)\n                log_once(\"fallback\")(f\"Fallback to raw DICOM for SID={sid}\")\n\n        if self.train:\n            X = _apply_tensor_aug(X)\n\n        return X, y\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T10:01:45.90823Z","iopub.execute_input":"2025-10-14T10:01:45.908472Z","iopub.status.idle":"2025-10-14T10:01:45.928096Z","shell.execute_reply.started":"2025-10-14T10:01:45.908446Z","shell.execute_reply":"2025-10-14T10:01:45.92756Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 5\ndef _build_backbone_safely():\n    \"\"\"\n    Prefer manual local load from the 'weights' dataset to avoid any download.\n    Fallback: random init if the file isn't present or load fails.\n    \"\"\"\n    local_ckpt = \"/kaggle/working/torch_home/hub/checkpoints/resnet18-f37072fd.pth\"\n\n    try:\n        if os.path.exists(local_ckpt):\n            bb_raw = resnet18(weights=None)  # build arch\n            state = torch.load(local_ckpt, map_location=\"cpu\")\n            # Load everything; we'll replace fc with Identity after\n            missing, unexpected = bb_raw.load_state_dict(state, strict=False)\n            bb_raw.fc = nn.Identity()\n            log(f\"Backbone: loaded local ImageNet weights. \"\n                f\"missing={len(missing)} unexpected={len(unexpected)}\")\n        else:\n            raise FileNotFoundError(local_ckpt)\n    except Exception as e:\n        log_warn(f\"Local pretrained load failed: {e} — using random init.\")\n        bb_raw = resnet18(weights=None)\n        bb_raw.fc = nn.Identity()\n\n    return nn.Sequential(bb_raw, nn.Flatten(1))  # outputs [N,512]\n\ndef _print_label_stats(df):\n    y = df[LABEL_COLS].values.astype(np.float32)\n    pos_rate = y.mean(axis=0)\n    ap_rate = pos_rate[-1]\n    log(f\"Label prevalence (mean across rows):\")\n    log(f\"  Aneurysm Present: {ap_rate:.4f}\")\n    for i in [0, 4, 8, 12]:\n        if i < len(LABEL_COLS)-1:\n            log(f\"  {LABEL_COLS[i]}: {pos_rate[i]:.4f}\")\n\ndef init_and_train_once():\n    global _initialized, model, AP_POS_WEIGHT\n    if _initialized:\n        log(\"init_and_train_once: already initialized.\")\n        return\n\n    set_seed(42)\n\n    df = pd.read_csv(TRAIN_CSV)\n    log(f\"Loaded train.csv: {len(df)} rows\")\n\n    if SUBSET_N and SUBSET_N < len(df):\n        df = df.sample(SUBSET_N, random_state=42)\n        log(f\"Subsampled to {len(df)} rows for runtime safety (SUBSET_N={SUBSET_N}).\")\n\n    tr_df, va_df = train_test_split(\n        df, test_size=0.15, random_state=42, stratify=df['Aneurysm Present']\n    )\n    log(f\"Split train/val: {len(tr_df)}/{len(va_df)}\")\n\n    _print_label_stats(tr_df)\n\n    pos = tr_df['Aneurysm Present'].sum()\n    neg = len(tr_df) - pos\n    AP_POS_WEIGHT = float(neg / max(1.0, pos))\n    log(f\"AP_POS_WEIGHT = {AP_POS_WEIGHT:.3f} (neg={int(neg)}, pos={int(pos)})\")\n\n    train_ds = RSNATensorDataset(tr_df, SERIES_ROOT, PREPROCESSED_ROOTS, k_slices=K_SLICES, img_size=IMG_SIZE, train=True)\n    val_ds   = RSNATensorDataset(va_df, SERIES_ROOT, PREPROCESSED_ROOTS, k_slices=K_SLICES, img_size=IMG_SIZE, train=False)\n\n    train_loader = DataLoader(train_ds, batch_size=BATCH_SZ, shuffle=True, num_workers=0, pin_memory=True)\n    val_loader   = DataLoader(val_ds,   batch_size=BATCH_SZ, shuffle=False, num_workers=0, pin_memory=True)\n    log(f\"Dataloaders ready — batch size {BATCH_SZ} | steps/epoch train={len(train_loader)} val={len(val_loader)}\")\n\n    backbone = _build_backbone_safely().to(device)\n    head = nn.Linear(512*2, 14).to(device)\n\n    col_w = COL_WEIGHTS.to(device)\n\n    def loss_fn(logits, targets):\n        pw = torch.ones(14, device=logits.device)\n        pw[-1] = AP_POS_WEIGHT if AP_POS_WEIGHT is not None else 1.0\n        bce = nn.functional.binary_cross_entropy_with_logits(\n            logits, targets, reduction='none', pos_weight=pw\n        )\n        bce = bce * col_w\n        return bce.mean()\n\n    optimizer = optim.AdamW(list(backbone.parameters())+list(head.parameters()),\n                            lr=2e-4, weight_decay=1e-4)\n    sched = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-5)\n\n    scaler = torch.amp.GradScaler('cuda', enabled=(device.type == \"cuda\"))\n    best_ap_auc = -1.0\n\n    for ep in range(EPOCHS):\n        backbone.train(); head.train(); tr_loss = 0.0\n        if torch.cuda.is_available():\n            torch.cuda.reset_peak_memory_stats()\n\n        for step, (X, y) in enumerate(train_loader, 1):\n            B,K,_,H,W = X.shape\n            X = X.view(B*K, 3, H, W).to(device, non_blocking=True)\n            y = y.to(device, non_blocking=True)\n\n            with torch.amp.autocast('cuda', enabled=(device.type == \"cuda\")):\n                feats = backbone(X)                    # [B*K,512]\n                feats = feats.view(B, K, -1)           # [B,K,512]\n                f_mean = feats.mean(dim=1)             # [B,512]\n                f_max  = feats.max(dim=1).values       # [B,512]\n                f_cat  = torch.cat([f_mean, f_max], dim=1)  # [B,1024]\n                logits = head(f_cat)                   # [B,14]\n                loss = loss_fn(logits, y)\n\n            optimizer.zero_grad(set_to_none=True)\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            tr_loss += loss.item() * B\n\n            if step % max(1, len(train_loader)//3) == 0 or step == 1:\n                lr = optimizer.param_groups[0]['lr']\n                if torch.cuda.is_available():\n                    mem = torch.cuda.memory_allocated()/1024**3\n                    peak = torch.cuda.max_memory_allocated()/1024**3\n                    log(f\"Epoch {ep+1}/{EPOCHS} step {step}/{len(train_loader)} | \"\n                        f\"loss {loss.item():.4f} | lr {lr:.2e} | mem {mem:.2f} GB (peak {peak:.2f})\")\n                else:\n                    log(f\"Epoch {ep+1}/{EPOCHS} step {step}/{len(train_loader)} | loss {loss.item():.4f} | lr {lr:.2e}\")\n\n        backbone.eval(); head.eval(); va_loss = 0.0\n        ap_y, ap_p = [], []\n        with torch.no_grad(), torch.amp.autocast('cuda', enabled=(device.type == \"cuda\")):\n            for X, y in val_loader:\n                B,K,_,H,W = X.shape\n                X = X.view(B*K, 3, H, W).to(device, non_blocking=True)\n                y = y.to(device, non_blocking=True)\n                feats = backbone(X).view(B, K, -1)\n                f_mean = feats.mean(dim=1)\n                f_max  = feats.max(dim=1).values\n                f_cat  = torch.cat([f_mean, f_max], dim=1)\n                logits = head(f_cat)\n                loss = loss_fn(logits, y)\n                va_loss += loss.item() * B\n                probs = torch.sigmoid(logits).detach().cpu().numpy()\n                ap_p.extend(probs[:, -1])\n                ap_y.extend(y[:, -1].detach().cpu().numpy().tolist())\n\n        try:\n            ap_auc = roc_auc_score(ap_y, ap_p)\n        except Exception as e:\n            log_warn(f\"roc_auc_score failed: {e}\")\n            ap_auc = float(\"nan\")\n\n        log(f\"Epoch {ep+1}/{EPOCHS} | Train {tr_loss/len(train_ds):.4f} | \"\n            f\"Val {va_loss/len(val_ds):.4f} | AneurysmPresent AUC {ap_auc:.4f}\")\n\n        if ap_auc > best_ap_auc:\n            best_ap_auc = ap_auc\n            try:\n                torch.save({\"backbone\": backbone.state_dict(),\n                            \"head\": head.state_dict()},\n                           \"/kaggle/working/rsna_resnet18_k9_best.pth\")\n                log(f\"Saved new best checkpoint with AP AUC {best_ap_auc:.4f}\")\n            except Exception as e:\n                log_warn(f\"Checkpoint save failed: {e}\")\n\n        sched.step()\n\n    class KFuseModel(nn.Module):\n        def __init__(self, backbone, head):\n            super().__init__()\n            self.backbone = backbone\n            self.head = head\n        def forward(self, xk):  # xk: [K,3,H,W]\n            K,_,H,W = xk.shape\n            feats = self.backbone(xk)      # [K,512]\n            f_mean = feats.mean(dim=0)     # [512]\n            f_max  = feats.max(dim=0).values\n            f_cat  = torch.cat([f_mean, f_max], dim=0)  # [1024]\n            return self.head(f_cat)        # [14]\n\n    model_fused = KFuseModel(backbone, head).to(device).eval()\n\n    try:\n        ck = torch.load(\"/kaggle/working/rsna_resnet18_k9_best.pth\", map_location=device)\n        model_fused.backbone.load_state_dict(ck[\"backbone\"])\n        model_fused.head.load_state_dict(ck[\"head\"])\n        log(\"Loaded best checkpoint into fused model.\")\n    except Exception as e:\n        log_warn(f\"Loading best checkpoint failed (continuing with last weights): {e}\")\n\n    globals()['model'] = model_fused\n    _initialized = True\n    log(\"Training init complete.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T10:01:45.928873Z","iopub.execute_input":"2025-10-14T10:01:45.92906Z","iopub.status.idle":"2025-10-14T10:01:45.955663Z","shell.execute_reply.started":"2025-10-14T10:01:45.929036Z","shell.execute_reply":"2025-10-14T10:01:45.955056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 6\ndef predict(series_path: str) -> pl.DataFrame | pd.DataFrame:\n    log(f\"predict() called for series_path={series_path}\")\n    init_and_train_once()\n\n    fpaths = []\n    for root, _, files in os.walk(series_path):\n        for f in files:\n            if f.lower().endswith(\".dcm\"):\n                fpaths.append(os.path.join(root, f))\n    fpaths.sort()\n    log(f\"Found {len(fpaths)} DICOM files for series.\")\n    if not fpaths:\n        raise FileNotFoundError(f\"No DICOMs in {series_path}\")\n\n    idxs = pick_k_indices(len(fpaths), K_SLICES)\n    imgs = [preprocess_dicom(fpaths[i], IMG_SIZE) for i in idxs]\n    X    = torch.stack(imgs, dim=0).float().clamp(0,1)  # [K,3,H,W]\n    log(f\"K-stack tensor shape: {tuple(X.shape)}\")\n\n    model.eval()\n    with torch.no_grad():\n        logits = model(X.to(device))\n        Xf = torch.flip(X, dims=[-1])\n        logits_f = model(Xf.to(device))\n\n        def rot_stack(X, deg):\n            out = []\n            for i in range(X.shape[0]):\n                out.append(TF.rotate(X[i], angle=deg, interpolation=InterpolationMode.BILINEAR))\n            return torch.stack(out, dim=0)\n        Xp = rot_stack(X, 7)\n        Xm = rot_stack(X, -7)\n        logits_p = model(Xp.to(device))\n        logits_m = model(Xm.to(device))\n\n        logits = (logits + logits_f + logits_p + logits_m) / 4.0\n        probs  = torch.sigmoid(logits).cpu().numpy()\n\n    series_id = os.path.basename(series_path.rstrip(\"/\"))\n    log(f\"Inference done for SID={series_id}; first 3 probs: {probs[:3]} … AP={probs[-1]:.4f}\")\n\n    predictions = pd.DataFrame([[series_id] + probs.tolist()], columns=[\"SeriesInstanceUID\", *LABEL_COLS])\n\n    try:\n        shutil.rmtree('/kaggle/shared', ignore_errors=True)\n    except Exception as e:\n        log_warn(f\"shared cleanup failed: {e}\")\n\n    return predictions.drop(columns=[\"SeriesInstanceUID\"])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T10:01:45.956322Z","iopub.execute_input":"2025-10-14T10:01:45.956509Z","iopub.status.idle":"2025-10-14T10:01:45.977975Z","shell.execute_reply.started":"2025-10-14T10:01:45.956493Z","shell.execute_reply":"2025-10-14T10:01:45.977426Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Cell 7\ninference_server = kaggle_evaluation.rsna_inference_server.RSNAInferenceServer(predict)\nlog(f\"IS_SUBMISSION={IS_SUBMISSION}\")\n\nif IS_SUBMISSION:\n    log(\"Starting server (submission mode).\")\n    inference_server.serve()\nelse:\n    log(\"Starting local gateway for preview.\")\n    shutil.rmtree(\"/kaggle/shared\", ignore_errors=True)\n    os.makedirs(\"/kaggle/shared\", exist_ok=True)\n    inference_server.run_local_gateway()\n\n    try:\n        pl.read_parquet(\"/kaggle/working/submission.parquet\").write_csv(\"/kaggle/working/submission.csv\")\n        log(\"Local preview written to /kaggle/working/submission.csv\")\n        display(pl.read_parquet(\"/kaggle/working/submission.parquet\").head())\n    except Exception as e:\n        log_warn(f\"No local preview: {e}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-14T10:01:45.978504Z","iopub.execute_input":"2025-10-14T10:01:45.978684Z","iopub.status.idle":"2025-10-14T10:21:49.405512Z","shell.execute_reply.started":"2025-10-14T10:01:45.97867Z","shell.execute_reply":"2025-10-14T10:21:49.404653Z"}},"outputs":[],"execution_count":null}]}