{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":"none","dataSources":[{"sourceType":"competition","sourceId":99552,"databundleVersionId":13441085},{"sourceType":"datasetVersion","sourceId":12850943,"datasetId":8128002,"databundleVersionId":13488655},{"sourceType":"kernelVersion","sourceId":257842109}],"dockerImageVersionId":31089,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =========================\n# Imports & Device Settings\n# =========================\nimport os\nimport gc\nimport shutil\nfrom collections import OrderedDict\nfrom typing import Tuple, List\n\nimport numpy as np\nimport polars as pl\nimport pydicom\nfrom scipy import ndimage\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch import amp\n\nimport kaggle_evaluation.rsna_inference_server\n\n# Device / AMP\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nUSE_AMP = torch.cuda.is_available()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-08-24T11:46:41.189443Z","iopub.execute_input":"2025-08-24T11:46:41.189815Z","iopub.status.idle":"2025-08-24T11:46:52.705391Z","shell.execute_reply.started":"2025-08-24T11:46:41.189784Z","shell.execute_reply":"2025-08-24T11:46:52.703993Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =====================\n# Competition Constants\n# =====================\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\n# =========\n# Filepaths\n# =========\nMODEL_WEIGHTS_PATH = \"/kaggle/input/rsna-train-voxel-medicalnet-resnet18-v1/model_weights.pth\"\n\n# ==========================\n# Processing / Model Configs\n# ==========================\nTARGET_SIZE = (64, 64, 64)      # (D,H,W) after final resize\nTARGET_SPACING_MM = 1.0         # isotropic spacing in mm\nCTA_WINDOW = (300.0, 700.0)     # (center, width) for CT\nMRI_Z_CLIP = 3.0                # z-score clip for MR\nLRU_CAPACITY = 8                # in-memory LRU capacity\n\n# =====================\n# MedicalNet モジュール読み込み\n# =====================\nimport sys\nsys.path.append(\"/kaggle/input/rsna2025-medicalnet-model/MedicalNet\")\nfrom models.resnet import resnet10, resnet18, resnet34, resnet50 \n\n# =====================\n# 使用するバックボーンの指定\n# =====================\nBACKBONE = 'resnet18'  # resnet10 / resnet18 / resnet34 / resnet50 から選択","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T11:46:52.707457Z","iopub.execute_input":"2025-08-24T11:46:52.708024Z","iopub.status.idle":"2025-08-24T11:46:52.746327Z","shell.execute_reply.started":"2025-08-24T11:46:52.707994Z","shell.execute_reply":"2025-08-24T11:46:52.744947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================\n# Utility Resizing Functions\n# ==========================\ndef _safe_zoom(volume: np.ndarray, zoom_factors: Tuple[float, ...], order: int = 1) -> np.ndarray:\n    \"\"\"Wrapper around ndimage.zoom with guards for invalid factors and shapes.\"\"\"\n    volume = np.nan_to_num(volume, copy=False)\n    zf = tuple(float(max(1e-6, f)) for f in zoom_factors)\n    if len(zf) != volume.ndim:\n        if len(zf) > volume.ndim:\n            zf = zf[:volume.ndim]\n        else:\n            zf = (1.0,) * (volume.ndim - len(zf)) + zf\n    return ndimage.zoom(volume, zf, order=order)\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) with bilinear order=1.\"\"\"\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T11:46:52.747511Z","iopub.execute_input":"2025-08-24T11:46:52.747905Z","iopub.status.idle":"2025-08-24T11:46:52.765309Z","shell.execute_reply.started":"2025-08-24T11:46:52.747839Z","shell.execute_reply":"2025-08-24T11:46:52.763936Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================\n# DICOM Series Processor\n# ==========================\nclass DICOMProcessor:\n    \"\"\"Turn a DICOM series folder into a normalized 3D volume (D,H,W) in [0,1].\"\"\"\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        lru_capacity: int = LRU_CAPACITY,\n    ):\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        self.memory_cache = OrderedDict()\n        self.lru_capacity = lru_capacity\n\n    # ---- LRU (memory only) ----\n    def _cache_put(self, key: str, vol: np.ndarray):\n        self.memory_cache[key] = vol\n        self.memory_cache.move_to_end(key)\n        if len(self.memory_cache) > self.lru_capacity:\n            self.memory_cache.popitem(last=False)\n\n    def _cache_get(self, key: str):\n        if key in self.memory_cache:\n            vol = self.memory_cache[key]\n            self.memory_cache.move_to_end(key)\n            return vol\n        return None\n\n    # ---- Public API ----\n    def load_dicom_series(self, series_path: str) -> np.ndarray:\n        \"\"\"Return (D,H,W) float32 in [0,1].\"\"\"\n        series_id = os.path.basename(series_path)\n\n        # Memory cache\n        m = self._cache_get(series_id)\n        if m is not None and isinstance(m, np.ndarray) and m.shape == self.target_size:\n            return m\n\n        try:\n            # Collect DICOM datasets\n            dicoms = []\n            for root, _, files in os.walk(series_path):\n                for f in files:\n                    if f.endswith(\".dcm\"):\n                        try:\n                            ds = pydicom.dcmread(os.path.join(root, f), force=True)\n                            if hasattr(ds, \"PixelData\"):\n                                dicoms.append(ds)\n                        except Exception as e:\n                            print(f\"[DICOM read] {e}\")\n                            continue\n            if not dicoms:\n                raise ValueError(f\"No valid DICOM files in {series_path}\")\n\n            # Sort slices by plane position along normal vector\n            dicoms = self._sort_slices(dicoms)\n\n            # Detect multiframe for spacing logic\n            has_multiframe = any(getattr(ds, \"NumberOfFrames\", 1) > 1 for ds in dicoms)\n\n            # Pixel spacing (dy,dx) and slice interval (dz)\n            spacing = self._get_spacing(dicoms, has_multiframe=has_multiframe)\n\n            # Choose base HxW without decoding entire stack twice\n            base_h, base_w = self._choose_base_shape(dicoms)\n\n            # Get modality tag\n            modality_tag = (getattr(dicoms[0], \"Modality\", \"\") or \"\").upper()\n\n            # Decode to slices (note: no per-frame reordering beyond dataset sorting)\n            vol_slices = []\n            for ds in dicoms:\n                arr = ds.pixel_array\n                if arr.ndim >= 3:\n                    h, w = arr.shape[-2], arr.shape[-1]\n                    n = int(np.prod(arr.shape[:-2]))\n                    frames = arr.reshape(n, h, w)\n                else:\n                    frames = arr[np.newaxis, ...]\n\n                for sl in frames:\n                    sl = sl.astype(np.float32)\n\n                    # Handle MONOCHROME1\n                    if getattr(ds, \"PhotometricInterpretation\", \"MONOCHROME2\") == \"MONOCHROME1\":\n                        sl = sl.max() - sl\n\n                    # Rescale\n                    slope = float(getattr(ds, \"RescaleSlope\", 1.0))\n                    intercept = float(getattr(ds, \"RescaleIntercept\", 0.0))\n                    sl = sl * slope + intercept\n\n                    # Resize 2D slice to base HxW\n                    sl = _resize_slice(sl, base_h, base_w)\n                    vol_slices.append(sl)\n\n            if len(vol_slices) == 0:\n                raise ValueError(\"No valid slices extracted\")\n\n            # Stack to (D,H,W)\n            volume = np.stack(vol_slices, axis=0).astype(np.float32)\n\n            # Modality-wise normalization to [0,1]\n            volume = self._normalize_by_modality(volume, modality_tag)\n\n            # Isotropic resample in mm\n            if self.target_spacing_mm is not None:\n                dz, dy, dx = spacing\n                z, y, x = volume.shape\n                newD = max(1, int(round(z * dz / self.target_spacing_mm)))\n                newH = max(1, int(round(y * dy / self.target_spacing_mm)))\n                newW = max(1, int(round(x * dx / self.target_spacing_mm)))\n                volume = _safe_zoom(volume, (newD / z, newH / y, newW / x), order=1)\n\n            # Final resize to target grid\n            tz, ty, tx = self.target_size\n            z, y, x = volume.shape\n            volume = _safe_zoom(volume, (tz / z, ty / y, tx / x), order=1).astype(np.float32)\n\n            # Put into cache and return\n            self._cache_put(series_id, volume)\n            return volume\n\n        except Exception as e:\n            print(f\"[Processor] Error: {e}\")\n            vol = np.zeros(self.target_size, dtype=np.float32)\n            self._cache_put(series_id, vol)\n            return vol\n\n    # ---- Helpers ----\n    def _sort_slices(self, ds_list: List[pydicom.dataset.FileDataset]) -> List[pydicom.dataset.FileDataset]:\n        \"\"\"Sort by dot(ImagePositionPatient, normal) or InstanceNumber.\"\"\"\n        try:\n            orient = np.array(ds_list[0].ImageOrientationPatient, dtype=np.float32)\n            row = orient[:3]; col = orient[3:]\n            normal = np.cross(row, col)\n            def sort_key(ds):\n                ipp = np.array(getattr(ds, \"ImagePositionPatient\", [0, 0, 0]), dtype=np.float32)\n                return float(np.dot(ipp, normal))\n            return sorted(ds_list, key=sort_key)\n        except Exception:\n            return sorted(ds_list, key=lambda ds: getattr(ds, \"InstanceNumber\", 0))\n\n    def _get_spacing(self, ds_sorted: List[pydicom.dataset.FileDataset], has_multiframe: bool = False) -> Tuple[float, float, float]:\n        \"\"\"Return (dz, dy, dx) in mm with robust fallbacks.\"\"\"\n        try:\n            dy, dx = map(float, ds_sorted[0].PixelSpacing)\n        except Exception:\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(getattr(ds_sorted[0], \"SpacingBetweenSlices\", getattr(ds_sorted[0], \"SliceThickness\", 1.0)))\n        else:\n            zs = []\n            for i in range(1, len(ds_sorted)):\n                p0 = np.array(getattr(ds_sorted[i-1], \"ImagePositionPatient\", [0, 0, 0]), dtype=np.float32)\n                p1 = np.array(getattr(ds_sorted[i], \"ImagePositionPatient\", [0, 0, 0]), dtype=np.float32)\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        \"\"\"Pick the most frequent (Rows, Columns) as base 2D shape.\"\"\"\n        shapes = []\n        for ds in ds_list:\n            try:\n                h, w = int(ds.Rows), int(ds.Columns)\n            except Exception:\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: window to [0,1]; MR: z-score -> clip -> [0,1].\"\"\"\n        volume = np.nan_to_num(volume, copy=False)\n        if modality_tag == \"CT\":\n            c, w = self.cta_window\n            lo, hi = c - w / 2.0, c + w / 2.0\n            v = np.clip(volume, lo, hi)\n            v = (v - lo) / (hi - lo + 1e-6)\n            return v.astype(np.float32, copy=False)\n        else:\n            mean = float(volume.mean())\n            std = float(volume.std() + 1e-6)\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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T11:46:52.766818Z","iopub.execute_input":"2025-08-24T11:46:52.767276Z","iopub.status.idle":"2025-08-24T11:46:52.804119Z","shell.execute_reply.started":"2025-08-24T11:46:52.76724Z","shell.execute_reply":"2025-08-24T11:46:52.802767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =====================\n# MedicalNet ResNet 推論用分類モデル\n# =====================\nclass MedicalNet3DClassifier(nn.Module):\n    def __init__(self, backbone_type='resnet18', num_classes=len(LABEL_COLS),\n                 weight_path=None, device='cuda'):\n        super().__init__()\n\n        # ResNet の種類選択\n        backbone_dict = {\n            'resnet10': resnet10,\n            'resnet18': resnet18,\n            'resnet34': resnet34,\n            'resnet50': resnet50\n        }\n        if backbone_type not in backbone_dict:\n            raise ValueError(f\"Invalid backbone_type: {backbone_type}\")\n        \n        backbone_fn = backbone_dict[backbone_type]\n\n        # backbone の生成（セグヘッド無効）\n        self.backbone = backbone_fn(\n            sample_input_W=TARGET_SIZE[1],\n            sample_input_H=TARGET_SIZE[0],\n            sample_input_D=MRI_Z_CLIP,\n            shortcut_type='B',\n            no_cuda=False,\n            num_seg_classes=0\n        )\n\n        # 最終チャネル数（ResNet の構造に依存）\n        channel_dict = {'resnet10': 512, 'resnet18': 512, 'resnet34': 512, 'resnet50': 2048}\n        in_features = channel_dict[backbone_type]\n\n        # 分類ヘッド\n        self.classifier = nn.Linear(in_features, num_classes)\n        state_dict = torch.load(weight_path, map_location=device)\n        self.load_state_dict(state_dict)\n        self.device = device\n        self.to(device)\n        self.eval()  # 推論用に eval モード\n\n    def forward(self, x):\n        # backbone の forward を使い conv_seg は無視\n        x = self.backbone.conv1(x)\n        x = self.backbone.bn1(x)\n        x = self.backbone.relu(x)\n        x = self.backbone.maxpool(x)\n        x = self.backbone.layer1(x)\n        x = self.backbone.layer2(x)\n        x = self.backbone.layer3(x)\n        x = self.backbone.layer4(x)\n\n        # Global Average Pooling\n        x = x.mean(dim=[2, 3, 4])\n        x = self.classifier(x)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T11:51:49.114377Z","iopub.execute_input":"2025-08-24T11:51:49.114771Z","iopub.status.idle":"2025-08-24T11:51:49.127228Z","shell.execute_reply.started":"2025-08-24T11:51:49.114731Z","shell.execute_reply":"2025-08-24T11:51:49.125714Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize processor (in-memory LRU only)\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    lru_capacity=LRU_CAPACITY,\n)\n\n# =====================\n# Create MedicalNet3DClassifier\n# =====================\nmodel = MedicalNet3DClassifier(\n    backbone_type=BACKBONE,\n    num_classes=len(LABEL_COLS),\n    weight_path=MODEL_WEIGHTS_PATH,\n    device=DEVICE\n)\nmodel.eval()\n\n# ==============\n# Inference API\n# ==============\n@torch.no_grad()\ndef predict(series_path: str) -> pl.DataFrame:\n    \"\"\"Server calls this function. Assumes global `model` and `processor` are ready.\"\"\"\n    try:\n        # CPU preprocessing\n        volume = processor.load_dicom_series(series_path)  # (D,H,W) in [0,1]\n        volume_tensor = torch.from_numpy(volume).unsqueeze(0).unsqueeze(0).to(DEVICE)  # (1,1,D,H,W)\n\n        # Forward pass\n        with amp.autocast(device_type='cuda', enabled=USE_AMP):\n            logits = model(volume_tensor)               # (1,14)\n            probs = torch.sigmoid(logits).cpu().numpy().flatten().tolist()\n\n        result_df = pl.DataFrame(data=[probs], schema=LABEL_COLS, orient='row')\n\n        # Cleanup\n        del volume_tensor\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        gc.collect()\n\n    except Exception as e:\n        print(f\"[Predict] Error: {e}\")\n        result_df = pl.DataFrame(data=[[0.5] * len(LABEL_COLS)], schema=LABEL_COLS, orient='row')\n\n    # Remove shared temp if present\n    shutil.rmtree('/kaggle/shared', ignore_errors=True)\n    return result_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T11:51:51.424368Z","iopub.execute_input":"2025-08-24T11:51:51.425315Z","iopub.status.idle":"2025-08-24T11:51:52.162534Z","shell.execute_reply.started":"2025-08-24T11:51:51.425265Z","shell.execute_reply":"2025-08-24T11:51:52.161477Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==========================\n# Run the Evaluation 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    try:\n        sub = pl.read_parquet('/kaggle/working/submission.parquet')\n        print(sub.head())\n    except Exception as e:\n        print(f\"Submission parquet not found yet: {e}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-24T11:49:15.349897Z","iopub.execute_input":"2025-08-24T11:49:15.350312Z","iopub.status.idle":"2025-08-24T11:49:52.153347Z","shell.execute_reply.started":"2025-08-24T11:49:15.350283Z","shell.execute_reply":"2025-08-24T11:49:52.151778Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}