{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"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"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":99552,"databundleVersionId":13851420,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13345699,"sourceType":"datasetVersion","datasetId":8461869},{"sourceId":13346023,"sourceType":"datasetVersion","datasetId":8453788,"isSourceIdPinned":true}],"dockerImageVersionId":31153,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"\n# RSNA Submission v1\n\nThis notebook prepares the trained `v1` model defined in `configs/submission_model/v1.yaml` for Kaggle inference. Update the configuration cell with your checkpoint path, run the notebook top to bottom, and it will launch the competition inference server to produce `submission.parquet`.\n","metadata":{}},{"cell_type":"code","source":"!rm /kaggle/working/* -rf","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-12T07:29:56.834025Z","iopub.execute_input":"2025-10-12T07:29:56.834307Z","iopub.status.idle":"2025-10-12T07:29:56.952623Z","shell.execute_reply.started":"2025-10-12T07:29:56.834281Z","shell.execute_reply":"2025-10-12T07:29:56.951434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-index --no-deps --find-links=/kaggle/input/rsna-submission-v1-wheels-data/wheelhouse hydra-core==1.3.2 monai==1.5.1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-12T07:29:56.954439Z","iopub.execute_input":"2025-10-12T07:29:56.954766Z","iopub.status.idle":"2025-10-12T07:30:01.299584Z","shell.execute_reply.started":"2025-10-12T07:29:56.954744Z","shell.execute_reply":"2025-10-12T07:30:01.298789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport gc\nimport logging\nimport os\nimport sys\nfrom functools import lru_cache\nfrom pathlib import Path\nfrom typing import Iterable, Sequence\n\nimport numpy as np\nimport polars as pl\nimport pydicom\nfrom hydra.utils import instantiate\nfrom omegaconf import OmegaConf\nfrom pydicom.dataset import FileDataset\nfrom pydicom.errors import InvalidDicomError\nfrom scipy import ndimage\nimport torch\n\nLOGGER = logging.getLogger(\"submission_v1\")\nif not LOGGER.handlers:\n    handler = logging.StreamHandler(sys.stdout)\n    handler.setFormatter(logging.Formatter(\"%(asctime)s - %(levelname)s - %(message)s\"))\n    LOGGER.addHandler(handler)\nLOGGER.setLevel(logging.INFO)\nLOGGER.propagate = False\n\nnp.random.seed(42)\ntorch.manual_seed(42)\ntorch.set_grad_enabled(False)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nLOGGER.info(\"Using device: %s\", device)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-12T07:30:01.300769Z","iopub.execute_input":"2025-10-12T07:30:01.301028Z","iopub.status.idle":"2025-10-12T07:30:06.370314Z","shell.execute_reply.started":"2025-10-12T07:30:01.301006Z","shell.execute_reply":"2025-10-12T07:30:06.369686Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# --- User configuration: update these paths for your Kaggle notebook run --- #\nREPO_ROOT = Path(\"/kaggle/input/kaggle-rsna2025/Kaggle-RSNA2025\")  # Update if the repository lives elsewhere\nif not REPO_ROOT.exists():\n    candidate_roots = [\n        Path(\"/kaggle/input/kaggle-rsna2025\"),\n        Path.cwd(),\n        Path.cwd().parent,\n    ]\n    for candidate in candidate_roots:\n        candidate = candidate.resolve()\n        if (candidate / \"configs/submission_model/v1.yaml\").exists():\n            REPO_ROOT = candidate\n            break\n\nCONFIG_PATH = REPO_ROOT / \"configs/submission_model/v1.yaml\"\nCHECKPOINT_PATH = Path(\"/kaggle/input/kaggle-rsna2025/Kaggle-RSNA2025/weights/best_model.pth\")\n\nif not CONFIG_PATH.exists():\n    raise FileNotFoundError(f\"Config file not found: {CONFIG_PATH}. Update REPO_ROOT.\")\nif not CHECKPOINT_PATH.exists():\n    raise FileNotFoundError(\"Checkpoint not found. Update CHECKPOINT_PATH to your trained weights.\")\n\nsys.path.append(str(REPO_ROOT))\nLOGGER.info(\"Repository root: %s\", REPO_ROOT)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-12T07:30:06.371797Z","iopub.execute_input":"2025-10-12T07:30:06.372084Z","iopub.status.idle":"2025-10-12T07:30:06.384347Z","shell.execute_reply.started":"2025-10-12T07:30:06.372066Z","shell.execute_reply":"2025-10-12T07:30:06.383702Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nfrom utils.builders import build_cls_model, build_data_aug\n\ncfg = OmegaConf.load(CONFIG_PATH)\nlabel_cols = list(cfg.training.dataset.label_cols)\nID_COL = cfg.training.dataset.id_col\nLOGGER.info(\"Loaded config with %d label columns\", len(label_cols))\n\nval_aug_mapping = build_data_aug(cfg.validation.data_aug)\nval_aug = val_aug_mapping.get(\"augmentation\")\nmodel = build_cls_model(cfg.model.backbone, cfg.model.cls_head)\n\nstate = torch.load(CHECKPOINT_PATH, map_location=\"cpu\")\nif isinstance(state, dict) and any(k.startswith(\"state_dict\") for k in state.keys()):\n    state = state.get(\"state_dict\", state.get(\"model_state_dict\"))\nmodel.load_state_dict(state, strict=True)\nmodel = model.to(device)\nmodel.eval()\nLOGGER.info(\"Model parameters loaded from %s\", CHECKPOINT_PATH)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-12T07:30:06.384954Z","iopub.execute_input":"2025-10-12T07:30:06.385164Z","iopub.status.idle":"2025-10-12T07:30:44.771865Z","shell.execute_reply.started":"2025-10-12T07:30:06.385148Z","shell.execute_reply":"2025-10-12T07:30:44.771149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# Processing / Model config\nTARGET_SIZE = (96, 96, 96)  # final (D,H,W)\nTARGET_SPACING_MM = 1.0  # isotropic resample\nCTA_WINDOW = (300.0, 700.0)  # (center, width) for CT (CTA)\nMRI_Z_CLIP = 3.0  # clip z-score to +/- 3 sigma\nSLOPE_MAX_ABS = 1000.0\nINTERCEPT_MAX_ABS = 10000.0\nRESCALE_MIN_LIMIT = -5000.0\nRESCALE_MAX_LIMIT = 10000.0\nCT_MIN_NORMAL = -2000.0\nCT_MAX_NORMAL = 4000.0\nMIN_WINDOW_WIDTH = 100.0\nSMALL_STD_EPS = 1e-6\nMULTIFRAME_DIM_THRESHOLD = 3\n\n\ndef _safe_zoom(volume: np.ndarray, zoom_factors: tuple[float, ...], order: int = 1) -> np.ndarray:\n    \"\"\"Robust wrapper around ndimage.zoom to avoid rank mismatch and invalid factors.\"\"\"\n    volume = np.nan_to_num(volume, copy=False)\n    zf = tuple(float(max(SMALL_STD_EPS, f)) for f in zoom_factors)  # avoid zeros/negatives\n    if len(zf) != volume.ndim:\n        zf = zf[: volume.ndim] if len(zf) > volume.ndim else (1.0,) * (volume.ndim - len(zf)) + zf\n    return ndimage.zoom(volume, zf, order=order)\n\n\ndef _resize_slice(arr: np.ndarray, out_h: int, out_w: int) -> np.ndarray:\n    \"\"\"Resize a 2D slice to (out_h, out_w) using safe zoom.\"\"\"\n    h, w = arr.shape\n    if h == out_h and w == out_w:\n        return arr.astype(np.float32, copy=False)\n    zy = out_h / max(h, 1)\n    zx = out_w / max(w, 1)\n    return _safe_zoom(arr, (zy, zx), order=1).astype(np.float32, copy=False)\n\n\nclass DICOMProcessor:\n    \"\"\"Process DICOM series into normalized 3D volumes.\"\"\"\n\n    def __init__(\n        self,\n        target_size: tuple[int, int, int] = TARGET_SIZE,\n        target_spacing_mm: float = TARGET_SPACING_MM,\n        cta_window: tuple[float, float] = CTA_WINDOW,\n        mri_z_clip: float = MRI_Z_CLIP,\n    ) -> None:\n        \"\"\"Initialize resampling, windowing, and normalization parameters.\"\"\"\n        self.target_size = target_size\n        self.target_spacing_mm = target_spacing_mm\n        self.cta_window = cta_window\n        self.mri_z_clip = mri_z_clip\n\n        # Adjustment counters\n        self.slope_adjustments = 0\n        self.intercept_adjustments = 0\n        self.adaptive_windowing_count = 0\n\n    def _validate_and_apply_rescale(self, sl: np.ndarray, ds: FileDataset) -> np.ndarray:\n        \"\"\"Validate slope/intercept values and apply rescaling.\"\"\"\n        slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n        intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n\n        # Validate slope\n        if slope <= 0 or not np.isfinite(slope) or abs(slope) > SLOPE_MAX_ABS:\n            slope = 1.0\n            self.slope_adjustments += 1\n\n        # Validate intercept with range clamping for extreme values\n        if not np.isfinite(intercept):\n            intercept = 0.0\n            self.intercept_adjustments += 1\n        elif abs(intercept) > INTERCEPT_MAX_ABS:\n            intercept = np.clip(intercept, CT_MIN_NORMAL, 0)\n            self.intercept_adjustments += 1\n\n        # Apply rescaling\n        rescaled = sl * slope + intercept\n\n        # Validate result\n        if np.any(~np.isfinite(rescaled)):\n            rescaled = np.nan_to_num(rescaled, copy=False)\n\n        # Post-rescale range check\n        min_val, max_val = rescaled.min(), rescaled.max()\n        if min_val < RESCALE_MIN_LIMIT or max_val > RESCALE_MAX_LIMIT:\n            rescaled = np.clip(rescaled, RESCALE_MIN_LIMIT * 0.6, RESCALE_MAX_LIMIT * 0.5)\n\n        return rescaled\n\n    def log_adjustment_summary(self) -> None:\n        \"\"\"Log summary of adjustments made during processing.\"\"\"\n        LOGGER.info(\n            \"Processing adjustments - Slope: %s, Intercept: %s, Adaptive windowing: %s\",\n            self.slope_adjustments,\n            self.intercept_adjustments,\n            self.adaptive_windowing_count,\n        )\n\n    def load_dicom_series(self, series_path: Path | str) -> np.ndarray:\n        \"\"\"Return (D,H,W) float32 volume in [0,1].\"\"\"\n        try:\n            return self._load_dicom_series(Path(series_path))\n        except (OSError, RuntimeError, ValueError):\n            LOGGER.warning(\"Failed to load series %s\", series_path, exc_info=True)\n            return np.zeros(self.target_size, dtype=np.float32)\n\n    def _load_dicom_series(self, path_obj: Path) -> np.ndarray:\n        dicoms = self._collect_dicoms(path_obj)\n        dicoms = self._sort_slices(dicoms)\n        modality_tag = (getattr(dicoms[0], \"Modality\", \"\") or \"\").upper()\n        has_multiframe = any(getattr(ds, \"NumberOfFrames\", 1) > 1 for ds in dicoms)\n        spacing = self._get_spacing(dicoms, has_multiframe=has_multiframe)\n\n        base_h, base_w = self._choose_base_shape(dicoms)\n        vol_slices = self._build_slices(dicoms, base_h, base_w)\n        if not vol_slices:\n            message = \"No valid slices extracted.\"\n            raise ValueError(message) from None\n\n        volume = np.stack(vol_slices, axis=0).astype(np.float32)\n        volume = self._normalize_by_modality(volume, modality_tag)\n\n        if self.target_spacing_mm is not None:\n            dz, dy, dx = spacing\n            z, y, x = volume.shape\n            new_d = max(1, round(z * dz / self.target_spacing_mm))\n            new_h = max(1, round(y * dy / self.target_spacing_mm))\n            new_w = max(1, round(x * dx / self.target_spacing_mm))\n            volume = _safe_zoom(volume, (new_d / z, new_h / y, new_w / x), order=1)\n\n        tz, ty, tx = self.target_size\n        z, y, x = volume.shape\n        return _safe_zoom(volume, (tz / z, ty / y, tx / x), order=1).astype(np.float32)\n\n    def _collect_dicoms(self, path_obj: Path) -> list[FileDataset]:\n        dicoms: list[FileDataset] = []\n        for dicom_path in path_obj.rglob(\"*.dcm\"):\n            try:\n                ds = pydicom.dcmread(dicom_path, force=True)\n            except (InvalidDicomError, OSError, ValueError):\n                LOGGER.debug(\"Skipping unreadable DICOM %s\", dicom_path, exc_info=LOGGER.isEnabledFor(logging.DEBUG))\n                continue\n            if hasattr(ds, \"PixelData\"):\n                dicoms.append(ds)\n        if not dicoms:\n            message = f\"No valid DICOM files with pixel data in {path_obj}\"\n            raise ValueError(message) from None\n        return dicoms\n\n    def _build_slices(\n        self,\n        dicoms: list[FileDataset],\n        base_h: int,\n        base_w: int,\n    ) -> list[np.ndarray]:\n        slices: list[np.ndarray] = []\n        for ds in dicoms:\n            arr = ds.pixel_array\n            if arr.ndim >= MULTIFRAME_DIM_THRESHOLD:\n                h, w = arr.shape[-2], arr.shape[-1]\n                frame_count = int(np.prod(arr.shape[:-2]))\n                frames = arr.reshape(frame_count, h, w)\n            else:\n                frames = arr[np.newaxis, ...]\n\n            for frame in frames:\n                slice_data = frame.astype(np.float32)\n                if getattr(ds, \"PhotometricInterpretation\", \"MONOCHROME2\") == \"MONOCHROME1\":\n                    slice_data = slice_data.max() - slice_data\n\n                slice_data = self._validate_and_apply_rescale(slice_data, ds)\n                slices.append(_resize_slice(slice_data, base_h, base_w))\n        return slices\n\n    def _sort_slices(self, ds_list: list[pydicom.dataset.FileDataset]) -> list[pydicom.dataset.FileDataset]:\n        try:\n            orient = np.array(ds_list[0].ImageOrientationPatient, dtype=np.float32)\n            row = orient[:3]\n            col = orient[3:]\n            normal = np.cross(row, col)\n\n            def sort_key(ds: FileDataset) -> float:\n                ipp = np.array(getattr(ds, \"ImagePositionPatient\", [0, 0, 0]), dtype=np.float32)\n                return float(np.dot(ipp, normal))\n\n            return sorted(ds_list, key=sort_key)\n        except (AttributeError, TypeError) as exc:\n            debug_context = LOGGER.isEnabledFor(logging.DEBUG)\n            LOGGER.debug(\"Falling back to instance number sort due to %s\", exc, exc_info=debug_context)\n            return sorted(ds_list, key=lambda ds: getattr(ds, \"InstanceNumber\", 0))\n\n    def _get_spacing(\n        self, ds_sorted: list[pydicom.dataset.FileDataset], *, has_multiframe: bool = False\n    ) -> tuple[float, float, float]:\n        try:\n            dy, dx = map(float, ds_sorted[0].PixelSpacing)\n        except (AttributeError, TypeError, ValueError) as exc:\n            debug_context = LOGGER.isEnabledFor(logging.DEBUG)\n            LOGGER.debug(\"Using default pixel spacing for %s due to %s\", ds_sorted[0], exc, exc_info=debug_context)\n            ps = getattr(ds_sorted[0], \"PixelSpacing\", [1.0, 1.0])\n            dy, dx = float(ps[0]), float(ps[1])\n\n        if has_multiframe:\n            dz = float(\n                getattr(\n                    ds_sorted[0],\n                    \"SpacingBetweenSlices\",\n                    getattr(ds_sorted[0], \"SliceThickness\", 1.0),\n                )\n            )\n        else:\n            zs = []\n            for i in range(1, len(ds_sorted)):\n                p0 = np.array(\n                    getattr(ds_sorted[i - 1], \"ImagePositionPatient\", [0, 0, 0]),\n                    dtype=np.float32,\n                )\n                p1 = np.array(\n                    getattr(ds_sorted[i], \"ImagePositionPatient\", [0, 0, 0]),\n                    dtype=np.float32,\n                )\n                d = np.linalg.norm(p1 - p0)\n                if d > 0:\n                    zs.append(d)\n            dz = float(np.median(zs)) if zs else float(getattr(ds_sorted[0], \"SliceThickness\", 1.0))\n\n        dz = dz if (dz > 0 and np.isfinite(dz)) else 1.0\n        dy = dy if (dy > 0 and np.isfinite(dy)) else 1.0\n        dx = dx if (dx > 0 and np.isfinite(dx)) else 1.0\n        return (dz, dy, dx)\n\n    def _choose_base_shape(self, ds_list: list[pydicom.dataset.FileDataset]) -> tuple[int, int]:\n        shapes = []\n        for ds in ds_list:\n            try:\n                h, w = int(ds.Rows), int(ds.Columns)\n            except (AttributeError, TypeError, ValueError):\n                arr = ds.pixel_array\n                h, w = arr.shape[-2], arr.shape[-1]\n            shapes.append((h, w))\n        vals, counts = np.unique(shapes, return_counts=True, axis=0)\n        base = tuple(vals[counts.argmax()])\n        return int(base[0]), int(base[1])\n\n    def _normalize_by_modality(self, volume: np.ndarray, modality_tag: str) -> np.ndarray:\n        \"\"\"CT: adaptive windowing for extreme ranges; MR: z-score -> clip -> [0,1].\"\"\"\n        volume = np.nan_to_num(volume, copy=False)\n\n        if modality_tag == \"CT\":\n            min_val, max_val = volume.min(), volume.max()\n\n            if min_val >= CT_MIN_NORMAL and max_val <= CT_MAX_NORMAL:\n                c, w = self.cta_window\n                lo, hi = c - w / 2.0, c + w / 2.0\n            else:\n                self.adaptive_windowing_count += 1\n\n                p1, p99 = np.percentile(volume, [1, 99])\n                margin = (p99 - p1) * 0.1\n                lo = p1 - margin\n                hi = p99 + margin\n\n                if hi - lo < MIN_WINDOW_WIDTH:\n                    center = (hi + lo) / 2\n                    half_width = MIN_WINDOW_WIDTH / 2\n                    lo = center - half_width\n                    hi = center + half_width\n\n            v = np.clip(volume, lo, hi)\n            v = (v - lo) / (hi - lo + 1e-6)\n            return v.astype(np.float32, copy=False)\n        # MRI processing\n        mean = float(volume.mean())\n        std = float(volume.std() + 1e-6)\n\n        # Validate statistics\n        if std < SMALL_STD_EPS or not np.isfinite(mean) or not np.isfinite(std):\n            return np.full_like(volume, 0.5, dtype=np.float32)\n\n        # Check dynamic range\n        min_val, max_val = volume.min(), volume.max()\n        if max_val - min_val < SMALL_STD_EPS:\n            return np.full_like(volume, 0.5, dtype=np.float32)\n\n        v = (volume - mean) / std\n        zc = float(self.mri_z_clip)\n        v = np.clip(v, -zc, zc)\n        v = (v + zc) / (2.0 * zc)\n        return v.astype(np.float32, copy=False)\n\n\nprocessor = DICOMProcessor(\n    target_size=TARGET_SIZE,\n    target_spacing_mm=TARGET_SPACING_MM,\n    cta_window=CTA_WINDOW,\n    mri_z_clip=MRI_Z_CLIP,\n)\nLOGGER.info(\"DICOM processor ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-12T07:30:44.772727Z","iopub.execute_input":"2025-10-12T07:30:44.773467Z","iopub.status.idle":"2025-10-12T07:30:44.807215Z","shell.execute_reply.started":"2025-10-12T07:30:44.773444Z","shell.execute_reply":"2025-10-12T07:30:44.806608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\ndef _prepare_tensor(volume: np.ndarray) -> torch.Tensor:\n    tensor = torch.from_numpy(volume).unsqueeze(0)  # (C, D, H, W)\n    tensor = val_aug(tensor)\n    return tensor.unsqueeze(0).to(device, non_blocking=True)  # (1, C, D, H, W)\n\n\n@torch.no_grad()\ndef predict(series_path: str) -> pl.DataFrame:\n    series_id = Path(series_path).name\n    LOGGER.info(f\"predict: {series_id}\")\n    volume = processor.load_dicom_series(series_path)\n    inputs = _prepare_tensor(volume)\n    logits = model(inputs)\n    probs = torch.sigmoid(logits).cpu().numpy()[0]\n    row = {ID_COL: series_id}\n    row.update({label: float(prob) for label, prob in zip(label_cols, probs, strict=True)})\n    gc.collect()\n    return pl.DataFrame(row)\n\n\nLOGGER.info(\"Predict function is ready\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-12T07:30:44.807943Z","iopub.execute_input":"2025-10-12T07:30:44.808167Z","iopub.status.idle":"2025-10-12T07:30:44.858614Z","shell.execute_reply.started":"2025-10-12T07:30:44.80815Z","shell.execute_reply":"2025-10-12T07:30:44.857851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nimport kaggle_evaluation.rsna_inference_server\n\ninference_server = kaggle_evaluation.rsna_inference_server.RSNAInferenceServer(predict)\n\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    inference_server.serve()\nelse:\n    inference_server.run_local_gateway()\n    display(pl.read_parquet('/kaggle/working/submission.parquet'))\n","metadata":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-10-12T07:30:44.859369Z","iopub.execute_input":"2025-10-12T07:30:44.859537Z","iopub.status.idle":"2025-10-12T07:31:13.245613Z","shell.execute_reply.started":"2025-10-12T07:30:44.859524Z","shell.execute_reply":"2025-10-12T07:31:13.24497Z"}},"outputs":[],"execution_count":null}]}