{"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":"gpu","dataSources":[{"sourceType":"competition","sourceId":99552,"databundleVersionId":13441085},{"sourceType":"datasetVersion","sourceId":12780021,"datasetId":8079690,"databundleVersionId":13404554}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA EfficientNetV2-S 3D Ensemble Pipeline  \nThis notebook implements a complete workflow for the RSNA competition:  \n- **DICOM → 3D Volume** preprocessing (resized to `(32, 384, 384)`)  \n- **EfficientNetV2-S** backbone with 32-channel input  \n- **5-Fold Ensemble** for robust predictions  \n- **Inference Server** for Kaggle competition integration  \n","metadata":{}},{"cell_type":"code","source":"# ====================================\n# 1. Imports\n# ====================================\nimport os\nimport numpy as np\nimport pydicom\nimport cv2\nfrom pathlib import Path\nfrom typing import List, Tuple, Dict, Optional\nfrom scipy import ndimage\nimport gc\nimport warnings\nwarnings.filterwarnings('ignore')\n\n\n# ====================================\n# 2. Preprocessor Class\n# ====================================\nclass DICOMPreprocessorKaggle:\n    \"\"\"\n    Preprocessor for Kaggle DICOM series → 3D volume\n    Converts CT/MR DICOM files into normalized volumes of fixed shape\n    \"\"\"\n\n    def __init__(self, target_shape: Tuple[int, int, int] = (32, 384, 384)):\n        self.target_depth, self.target_height, self.target_width = target_shape\n\n    # -------------------------------\n    # Load series\n    # -------------------------------\n    def load_dicom_series(self, series_path: str) -> Tuple[List[pydicom.Dataset], str]:\n        series_path = Path(series_path)\n        series_name = series_path.name\n\n        dicom_files = []\n        for root, _, files in os.walk(series_path):\n            for file in files:\n                if file.endswith('.dcm'):\n                    dicom_files.append(os.path.join(root, file))\n\n        if not dicom_files:\n            raise ValueError(f\"No DICOM files found in {series_path}\")\n\n        datasets = []\n        for filepath in dicom_files:\n            try:\n                ds = pydicom.dcmread(filepath, force=True)\n                datasets.append(ds)\n            except:\n                continue\n\n        if not datasets:\n            raise ValueError(f\"No valid DICOM files in {series_path}\")\n\n        print(f\"Loaded {len(datasets)} files from series {series_name}\")\n        return datasets, series_name\n\n    # -------------------------------\n    # Slice handling\n    # -------------------------------\n    def extract_slice_info(self, datasets: List[pydicom.Dataset]) -> List[Dict]:\n        slice_info = []\n        for i, ds in enumerate(datasets):\n            info = {\n                'dataset': ds,\n                'index': i,\n                'instance_number': getattr(ds, 'InstanceNumber', i),\n            }\n            pos = getattr(ds, 'ImagePositionPatient', None)\n            info['z_position'] = float(pos[2]) if pos is not None else float(info['instance_number'])\n            slice_info.append(info)\n        return slice_info\n\n    def sort_slices_by_position(self, slice_info: List[Dict]) -> List[Dict]:\n        return sorted(slice_info, key=lambda x: x['z_position'])\n\n    # -------------------------------\n    # Windowing & normalization\n    # -------------------------------\n    def get_windowing_params(self, ds: pydicom.Dataset) -> Tuple[Optional[float], Optional[float]]:\n        modality = getattr(ds, 'Modality', 'CT')\n        if modality == 'CT':\n            return 50, 350  # CTA windowing defaults\n        elif modality == 'MR':\n            return None, None\n        else:\n            return None, None\n\n    def apply_windowing_or_normalize(self, img: np.ndarray, center: Optional[float], width: Optional[float]) -> np.ndarray:\n        if center is not None and width is not None:\n            # Percentile-based normalization for CT\n            p1, p99 = 0, 500\n            normalized = np.clip(img, p1, p99)\n            normalized = (normalized - p1) / (p99 - p1 + 1e-7)\n        else:\n            # Percentile-based normalization for MR\n            p1, p99 = np.percentile(img, [1, 99])\n            normalized = np.clip(img, p1, p99)\n            normalized = (normalized - p1) / (p99 - p1 + 1e-7)\n        return (normalized * 255).astype(np.uint8)\n\n    # -------------------------------\n    # Pixel extraction\n    # -------------------------------\n    def extract_pixel_array(self, ds: pydicom.Dataset) -> np.ndarray:\n        img = ds.pixel_array.astype(np.float32)\n        if img.ndim == 3:  # take middle frame if 3D inside 2D file\n            img = img[img.shape[0] // 2]\n        if img.ndim == 3 and img.shape[-1] == 3:  # convert RGB to grayscale\n            img = cv2.cvtColor(img.astype(np.uint8), cv2.COLOR_RGB2GRAY).astype(np.float32)\n        return img\n\n    # -------------------------------\n    # Resize volume\n    # -------------------------------\n    def resize_volume_3d(self, volume: np.ndarray) -> np.ndarray:\n        target_shape = (self.target_depth, self.target_height, self.target_width)\n        if volume.shape == target_shape:\n            return volume\n\n        zoom_factors = [target_shape[i] / volume.shape[i] for i in range(3)]\n        resized_volume = ndimage.zoom(volume, zoom_factors, order=1, mode='nearest')\n        resized_volume = resized_volume[:self.target_depth, :self.target_height, :self.target_width]\n\n        pad_width = [\n            (0, max(0, self.target_depth - resized_volume.shape[0])),\n            (0, max(0, self.target_height - resized_volume.shape[1])),\n            (0, max(0, self.target_width - resized_volume.shape[2]))\n        ]\n        if any(p[1] > 0 for p in pad_width):\n            resized_volume = np.pad(resized_volume, pad_width, mode='edge')\n\n        return resized_volume.astype(np.uint8)\n\n    # -------------------------------\n    # Main processing\n    # -------------------------------\n    def process_series(self, series_path: str) -> np.ndarray:\n        datasets, series_name = self.load_dicom_series(series_path)\n        first_ds = datasets[0]\n        first_img = first_ds.pixel_array\n\n        if len(datasets) == 1 and first_img.ndim == 3:\n            print(f\"Processing single 3D DICOM: {series_name}\")\n            return self._process_single_3d_dicom(first_ds)\n        else:\n            print(f\"Processing {len(datasets)} 2D DICOMs: {series_name}\")\n            return self._process_multiple_2d_dicoms(datasets)\n\n    def _process_single_3d_dicom(self, ds: pydicom.Dataset) -> np.ndarray:\n        volume = ds.pixel_array.astype(np.float32)\n        center, width = self.get_windowing_params(ds)\n        processed = [self.apply_windowing_or_normalize(slice_img, center, width) for slice_img in volume]\n        volume = np.stack(processed, axis=0)\n        return self.resize_volume_3d(volume)\n\n    def _process_multiple_2d_dicoms(self, datasets: List[pydicom.Dataset]) -> np.ndarray:\n        slice_info = self.extract_slice_info(datasets)\n        sorted_slices = self.sort_slices_by_position(slice_info)\n        first_img = self.extract_pixel_array(sorted_slices[0]['dataset'])\n        center, width = self.get_windowing_params(sorted_slices[0]['dataset'])\n\n        processed = []\n        for slice_data in sorted_slices:\n            img = self.extract_pixel_array(slice_data['dataset'])\n            img = self.apply_windowing_or_normalize(img, center, width)\n            img = cv2.resize(img, (self.target_width, self.target_height))\n            processed.append(img)\n\n        volume = np.stack(processed, axis=0)\n        return self.resize_volume_3d(volume)\n\n\n# ====================================\n# 3. Helper functions\n# ====================================\ndef process_dicom_series_kaggle(series_path: str, target_shape: Tuple[int, int, int] = (32, 384, 384)) -> np.ndarray:\n    preprocessor = DICOMPreprocessorKaggle(target_shape=target_shape)\n    return preprocessor.process_series(series_path)\n\n\ndef process_dicom_series_safe(series_path: str, target_shape: Tuple[int, int, int] = (32, 384, 384)) -> np.ndarray:\n    try:\n        return process_dicom_series_kaggle(series_path, target_shape)\n    finally:\n        gc.collect()\n\n\n# ====================================\n# 4. Test function\n# ====================================\ndef test_single_series(series_path: str, target_shape: Tuple[int, int, int] = (32, 384, 384)):\n    volume = process_dicom_series_safe(series_path, target_shape)\n    if volume is not None:\n        print(f\"✓ Processed: {series_path}\")\n        print(f\"  Shape: {volume.shape}, dtype: {volume.dtype}, range: [{volume.min()}, {volume.max()}]\")\n    else:\n        print(f\"✗ Failed to process {series_path}\")\n    return volume","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-09-06T19:42:34.392995Z","iopub.execute_input":"2025-09-06T19:42:34.393349Z","iopub.status.idle":"2025-09-06T19:42:34.416316Z","shell.execute_reply.started":"2025-09-06T19:42:34.393325Z","shell.execute_reply":"2025-09-06T19:42:34.415673Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# Imports\n# ====================================================\nimport sys\nimport gc\nimport json\nimport shutil\nimport warnings\nwarnings.filterwarnings('ignore')\n\nfrom pathlib import Path\nfrom typing import List, Dict, Optional, Tuple\n\n# Data handling\nimport numpy as np\nimport polars as pl\nimport pandas as pd\n\n# Medical imaging\nimport pydicom\nimport cv2\n\n# ML/DL\nimport torch\nimport torch.nn as nn\nfrom torch.cuda.amp import autocast\nimport timm\n\n# Transformations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# Competition API\nimport kaggle_evaluation.rsna_inference_server\n\n\n# ====================================================\n# Device setup\n# ====================================================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"[INFO] Using device: {device}\")\n\n\n# ====================================================\n# Competition constants\n# ====================================================\nID_COLUMN = \"SeriesInstanceUID\"\nLABEL_COLUMNS = [\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# ====================================================\n# Configuration\n# ====================================================\nclass InferenceConfig:\n    model_name = \"tf_efficientnetv2_s.in21k_ft_in1k\"\n    image_size = 384\n    num_classes = len(LABEL_COLUMNS)\n    input_channels = 32\n\n    # Volume preprocessing\n    target_shape = (32, 384, 384)\n\n    # Inference\n    batch_size = 1\n    use_amp = True\n    use_tta = False  # Left/right symmetry prohibits TTA\n    tta_transforms = 0\n\n    # Model files\n    model_dir = \"/kaggle/input/rsna2025-effnetv2-32ch\"\n    folds = [0, 1, 2, 3, 4, 5]\n\n    # Ensemble weights (None = equal weighting)\n    ensemble_weights = None\n\n\nConfig = InferenceConfig()\n\n\n# ====================================================\n# Transforms\n# ====================================================\ndef get_inference_transform():\n    return A.Compose([\n        A.Resize(Config.image_size, Config.image_size),\n        A.Normalize(),\n        ToTensorV2(),\n    ])\n\n\n# ====================================================\n# Model loading\n# ====================================================\nmodels_dict = {}\nbase_transform = None\n\ndef load_model_fold(fold: int) -> nn.Module:\n    \"\"\"Load one fold of the trained model\"\"\"\n    model_path = Path(Config.model_dir) / f\"{Config.model_name}_fold{fold}_best.pth\"\n    if not model_path.exists():\n        raise FileNotFoundError(f\"Missing model file: {model_path}\")\n\n    checkpoint = torch.load(model_path, map_location=device, weights_only=False)\n    model = timm.create_model(\n        Config.model_name,\n        num_classes=Config.num_classes,\n        pretrained=False,\n        in_chans=Config.input_channels\n    )\n    model.load_state_dict(checkpoint[\"model\"])\n    model = model.to(device).eval()\n    return model\n\n\ndef load_all_models():\n    \"\"\"Load models for all folds\"\"\"\n    global models_dict, base_transform\n\n    print(\"[INFO] Loading models...\")\n    for fold in Config.folds:\n        try:\n            models_dict[fold] = load_model_fold(fold)\n            print(f\"[INFO] Loaded fold {fold}\")\n        except Exception as e:\n            print(f\"[WARN] Could not load fold {fold}: {e}\")\n\n    if not models_dict:\n        raise ValueError(\"No models were successfully loaded\")\n\n    base_transform = get_inference_transform()\n\n    # Warm-up pass\n    dummy_input = torch.randn(1, Config.input_channels, Config.image_size, Config.image_size).to(device)\n    with torch.no_grad():\n        for f, m in models_dict.items():\n            _ = m(dummy_input)\n\n    print(f\"[INFO] Loaded {len(models_dict)} models: {list(models_dict.keys())}\")\n\n\n# ====================================================\n# Prediction functions\n# ====================================================\ndef predict_with_model(model: nn.Module, volume: np.ndarray) -> np.ndarray:\n    \"\"\"Predict with a single model\"\"\"\n    # Input volume: (D,H,W) -> (H,W,D)\n    volume = volume.transpose(1, 2, 0)\n\n    transformed = base_transform(image=volume)\n    tensor = transformed[\"image\"].unsqueeze(0).to(device)\n\n    with torch.no_grad():\n        with autocast(enabled=Config.use_amp):\n            output = model(tensor)\n            return torch.sigmoid(output).cpu().numpy().squeeze()\n\n\ndef predict_with_ensemble(volume: np.ndarray) -> np.ndarray:\n    \"\"\"Average predictions across folds\"\"\"\n    preds, weights = [], []\n\n    for fold, model in models_dict.items():\n        preds.append(predict_with_model(model, volume))\n        if Config.ensemble_weights:\n            weights.append(Config.ensemble_weights.get(fold, 1.0))\n        else:\n            weights.append(1.0)\n\n    weights = np.array(weights) / np.sum(weights)\n    preds = np.array(preds)\n    return np.average(preds, weights=weights, axis=0)\n\n\n# ====================================================\n# Series prediction\n# ====================================================\ndef _predict_series(series_path: str) -> pl.DataFrame:\n    \"\"\"Core logic for inference on a single DICOM series\"\"\"\n    if not models_dict:\n        load_all_models()\n\n    series_id = os.path.basename(series_path)\n\n    try:\n        preprocessor = DICOMPreprocessorKaggle(target_shape=Config.target_shape)\n        volume = preprocessor.process_series(series_path)\n\n        final_pred = predict_with_ensemble(volume)\n        dfs = pl.DataFrame(\n            data=[[series_id] + final_pred.tolist()],\n            schema=[ID_COLUMN] + LABEL_COLUMNS,\n            orient=\"row\"\n        )\n        return dfs.drop(ID_COLUMN)\n\n    except Exception as e:\n        print(f\"[ERROR] Series {series_id} failed: {e}\")\n        fallback = [0.1] * len(LABEL_COLUMNS)\n        return pl.DataFrame([fallback], schema=LABEL_COLUMNS)\n\n\n# ====================================================\n# Public API functions\n# ====================================================\ndef predict(series_path: str) -> pl.DataFrame:\n    \"\"\"Main prediction entry point (used by server)\"\"\"\n    try:\n        return _predict_series(series_path)\n    except Exception as e:\n        print(f\"[ERROR] Prediction failed for {os.path.basename(series_path)}: {e}\")\n        fallback = [0.1] * len(LABEL_COLUMNS)\n        return pl.DataFrame([fallback], schema=LABEL_COLUMNS)\n    finally:\n        shared_dir_path = \"/kaggle/shared\"\n        shutil.rmtree(shared_dir_path, ignore_errors=True)\n        os.makedirs(shared_dir_path, exist_ok=True)\n\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        gc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-06T19:42:34.528077Z","iopub.execute_input":"2025-09-06T19:42:34.528773Z","iopub.status.idle":"2025-09-06T19:42:34.561486Z","shell.execute_reply.started":"2025-09-06T19:42:34.528746Z","shell.execute_reply":"2025-09-06T19:42:34.560621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# Entry Point\n# ====================================================\n\nprint(\"[INFO] Starting setup...\")\n\n# Preload all trained folds\nload_all_models()\n\n# Initialize RSNA server\nserver = kaggle_evaluation.rsna_inference_server.RSNAInferenceServer(predict)\n\n# Decide execution mode (competition vs local)\nif os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n    print(\"[INFO] Running in COMPETITION mode\")\n    server.serve()\nelse:\n    print(\"[INFO] Running in LOCAL mode\")\n    server.run_local_gateway()\n\n    # Load and show a sample of the submission\n    sub_df = pl.read_parquet(\"/kaggle/working/submission.parquet\")\n    print(f\"[INFO] Submission shape: {sub_df.shape}\")\n    display(sub_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-06T19:42:34.562488Z","iopub.execute_input":"2025-09-06T19:42:34.562769Z","iopub.status.idle":"2025-09-06T19:42:47.229753Z","shell.execute_reply.started":"2025-09-06T19:42:34.56275Z","shell.execute_reply":"2025-09-06T19:42:47.22914Z"}},"outputs":[],"execution_count":null}]}