{"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":[{"sourceId":99552,"databundleVersionId":13762876,"sourceType":"competition"},{"sourceId":12687919,"sourceType":"datasetVersion","datasetId":7976292},{"sourceId":12780021,"sourceType":"datasetVersion","datasetId":8079690}],"dockerImageVersionId":31090,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Setup and Imports","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport gc\nimport json\nimport shutil\nimport warnings\nfrom pathlib import Path\nfrom typing import List, Dict, Tuple\n\nimport numpy as np\nimport polars as pl\nimport pandas as pd\nimport pydicom\nimport cv2\nfrom scipy import ndimage\n\nimport torch\nimport torch.nn as nn\nfrom torch.cuda.amp import autocast\nimport timm\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport kaggle_evaluation.rsna_inference_server\n\nwarnings.filterwarnings('ignore')\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-08T11:31:37.115309Z","iopub.execute_input":"2025-09-08T11:31:37.115978Z","iopub.status.idle":"2025-09-08T11:32:22.196407Z","shell.execute_reply.started":"2025-09-08T11:31:37.11594Z","shell.execute_reply":"2025-09-08T11:32:22.195786Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Configuration","metadata":{}},{"cell_type":"code","source":"class CFG:\n    # Competition constants\n    ID_COL = 'SeriesInstanceUID'\n    LABEL_COLS = [\n        'Left Infraclinoid Internal Carotid Artery', 'Right Infraclinoid Internal Carotid Artery',\n        'Left Supraclinoid Internal Carotid Artery', 'Right Supraclinoid Internal Carotid Artery',\n        'Left Middle Cerebral Artery', 'Right Middle Cerebral Artery',\n        'Anterior Communicating Artery', 'Left Anterior Cerebral Artery', 'Right Anterior Cerebral Artery',\n        'Left Posterior Communicating Artery', 'Right Posterior Communicating Artery',\n        'Basilar Tip', 'Other Posterior Circulation', 'Aneurysm Present',\n    ]\n\n    # --- Pipeline 1: 3D Voxel Model Config ---\n    P1_MODEL_DIR = '/kaggle/input/rsna2025-effnetv2-32ch/'\n    P1_MODEL_NAME = 'tf_efficientnetv2_s.in21k_ft_in1k'\n    P1_MODEL_FOLDS = [0, 1, 2, 3, 4]\n    P1_INPUT_SHAPE = (32, 256, 256)  # D, H, W\n    P1_IN_CHANS = 32\n\n    # --- Pipeline 2: 2.5D Projection Model Config ---\n    P2_MODEL_DIR = '/kaggle/input/rsna-iad-trained-models/models/'\n    P2_INPUT_SIZE = 512\n    P2_MODELS = {\n        'tf_efficientnetv2_s': 'tf_efficientnetv2_s_fold0_best.pth',\n        'convnext_small': 'convnext_small_fold0_best.pth',\n        'swin_small_patch4_window7_224': 'swin_small_patch4_window7_224_fold0_best.pth',\n    }\n    P2_ENSEMBLE_WEIGHTS = {\n        'tf_efficientnetv2_s': 0.4,\n        'convnext_small': 0.3,\n        'swin_small_patch4_window7_224': 0.3,\n    }\n\n    # --- Final Ensemble Weights ---\n    PIPELINE_3D_WEIGHT = 0.5\n    PIPELINE_2D_WEIGHT = 0.5\n\n    # --- Inference Config ---\n    USE_TTA = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-08T11:32:22.197579Z","iopub.execute_input":"2025-09-08T11:32:22.197802Z","iopub.status.idle":"2025-09-08T11:32:22.202765Z","shell.execute_reply.started":"2025-09-08T11:32:22.197783Z","shell.execute_reply":"2025-09-08T11:32:22.202231Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Unified Preprocessing","metadata":{}},{"cell_type":"code","source":"def load_and_sort_dicom_series(series_path: str) -> List[pydicom.Dataset]:\n    \"\"\"\n    Loads, validates, and sorts DICOM files using IOP-aware sorting.\n    \"\"\"\n    series_path = Path(series_path)\n    # Read all files, not just .dcm, and filter by DICOM validity.\n    files = [os.path.join(r, f) for r, _, fs in os.walk(series_path) for f in fs if not f.startswith('.')]\n\n    datasets = []\n    for fp in files:\n        try:\n            ds = pydicom.dcmread(fp, force=True)\n            if 'PixelData' in ds:\n                datasets.append(ds)\n        except Exception:\n            continue\n\n    if not datasets:\n        raise ValueError(f\"No valid DICOM files could be read from {series_path}\")\n\n    def slice_pos(ds):\n        # IOP-aware projection of IPP onto slice normal for robust sorting\n        ipp = np.array(getattr(ds, 'ImagePositionPatient', [0, 0, 0]), dtype=float)\n        iop = getattr(ds, 'ImageOrientationPatient', None)\n        if iop is not None and len(iop) >= 6:\n            row = np.array(iop[:3], dtype=float)\n            col = np.array(iop[3:6], dtype=float)\n            normal = np.cross(row, col)\n            return float(np.dot(ipp, normal))\n        # Fallbacks if IOP is missing\n        if hasattr(ds, 'SliceLocation'):\n            return float(ds.SliceLocation)\n        return float(getattr(ds, 'InstanceNumber', 0))\n\n    datasets.sort(key=slice_pos)\n    return datasets\n\ndef process_to_volume(datasets: List[pydicom.Dataset]) -> Tuple[np.ndarray, Dict]:\n    \"\"\"Processes sorted DICOMs into a 3D volume and extracts metadata.\"\"\"\n    slices = []\n    metadata = {}\n\n    for i, ds in enumerate(datasets):\n        img = ds.pixel_array.astype(np.float32)\n\n        slope = float(getattr(ds, 'RescaleSlope', 1))\n        intercept = float(getattr(ds, 'RescaleIntercept', 0))\n        img = img * slope + intercept\n\n        # Use CTA windowing as a robust default\n        center, width = 50, 350\n        img_min = center - width / 2\n        img_max = center + width / 2\n        img = np.clip(img, img_min, img_max)\n\n        slices.append(img)\n\n        if i == 0:  # Extract metadata from the first slice\n            try:\n                age_str = getattr(ds, 'PatientAge', '050Y')\n                metadata['age'] = int(''.join(filter(str.isdigit, age_str[:3])) or '50')\n            except:\n                metadata['age'] = 50\n            metadata['sex'] = 1 if getattr(ds, 'PatientSex', 'M') == 'M' else 0\n\n    if not slices:\n        raise ValueError(\"Could not extract any pixel arrays from the DICOM series.\")\n\n    return np.stack(slices, axis=0), metadata\n\ndef unified_preprocessor(series_path: str) -> Tuple[np.ndarray, np.ndarray, Dict]:\n    \"\"\"\n    Main preprocessing function.\n    Returns:\n        - volume_3d (np.ndarray): (D, H, W) for the 3D pipeline.\n        - proj_image_2d (np.ndarray): (H, W, 3) for the 2.5D pipeline.\n        - metadata (Dict): Patient age and sex.\n    \"\"\"\n    datasets = load_and_sort_dicom_series(series_path)\n    full_volume, metadata = process_to_volume(datasets)\n\n    # --- 1. Create 3D Voxel Volume ---\n    target_d, target_h, target_w = CFG.P1_INPUT_SHAPE\n    zoom_factors = [\n        target_d / full_volume.shape[0],\n        target_h / full_volume.shape[1],\n        target_w / full_volume.shape[2]\n    ]\n    volume_3d = ndimage.zoom(full_volume, zoom_factors, order=1, mode='nearest')\n\n    # Normalize to [0, 255] uint8\n    vol_min, vol_max = volume_3d.min(), volume_3d.max()\n    if vol_max > vol_min:\n        volume_3d = ((volume_3d - vol_min) / (vol_max - vol_min) * 255).astype(np.uint8)\n    else:\n        volume_3d = np.zeros_like(volume_3d, dtype=np.uint8)\n\n    # --- 2. Create 2.5D Projection Image ---\n    size = CFG.P2_INPUT_SIZE\n\n    # Projections are created from the original full_volume for max quality\n    middle_slice = cv2.resize(full_volume[full_volume.shape[0] // 2], (size, size), interpolation=cv2.INTER_AREA)\n    mip = cv2.resize(np.max(full_volume, axis=0), (size, size), interpolation=cv2.INTER_AREA)\n    std_proj = cv2.resize(np.std(full_volume, axis=0), (size, size), interpolation=cv2.INTER_AREA)\n\n    # Normalize each channel to [0, 255] uint8\n    def normalize_channel(ch):\n        ch_min, ch_max = ch.min(), ch.max()\n        if ch_max > ch_min:\n            return ((ch - ch_min) / (ch_max - ch_min) * 255).astype(np.uint8)\n        return np.zeros_like(ch, dtype=np.uint8)\n\n    proj_image_2d = np.stack([\n        normalize_channel(middle_slice),\n        normalize_channel(mip),\n        normalize_channel(std_proj)\n    ], axis=-1)\n\n    return volume_3d, proj_image_2d, metadata","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-08T11:32:22.203757Z","iopub.execute_input":"2025-09-08T11:32:22.204085Z","iopub.status.idle":"2025-09-08T11:32:22.254828Z","shell.execute_reply.started":"2025-09-08T11:32:22.204056Z","shell.execute_reply":"2025-09-08T11:32:22.254103Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Model Definitions and Transforms","metadata":{}},{"cell_type":"code","source":"# --- Model Definitions ---\nclass Timm3DModel(nn.Module):\n    def __init__(self, model_name, pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(\n            model_name,\n            pretrained=pretrained,\n            in_chans=CFG.P1_IN_CHANS,\n            num_classes=len(CFG.LABEL_COLS)\n        )\n    def forward(self, x):\n        return self.model(x)\n\nclass MultiBackboneModel(nn.Module):\n    def __init__(self, model_name, num_classes=len(CFG.LABEL_COLS), pretrained=False):\n        super().__init__()\n        \n        create_kwargs = {\n            'pretrained': pretrained,\n            'num_classes': 0,\n            'in_chans': 3,\n            'global_pool': '',\n        }\n        if 'swin' in model_name:\n            create_kwargs['img_size'] = CFG.P2_INPUT_SIZE\n            \n        self.backbone = timm.create_model(model_name, **create_kwargs)\n        \n        with torch.no_grad():\n            dummy_features = self.backbone(torch.randn(1, 3, CFG.P2_INPUT_SIZE, CFG.P2_INPUT_SIZE))\n            if dummy_features.ndim == 4: # Conv features (N, C, H, W)\n                num_features = dummy_features.shape[1]\n                self.pool = nn.AdaptiveAvgPool2d(1)\n            else: # Transformer features (N, T, C)\n                num_features = dummy_features.shape[-1]\n                self.pool = lambda x: x.mean(dim=1)\n\n        self.meta_fc = nn.Sequential(\n            nn.Linear(2, 16), nn.ReLU(), nn.Dropout(0.2), nn.Linear(16, 32), nn.ReLU()\n        )\n\n        self.classifier = nn.Sequential(\n            nn.Linear(num_features + 32, 512), nn.BatchNorm1d(512), nn.ReLU(), nn.Dropout(0.3),\n            nn.Linear(512, 256), nn.BatchNorm1d(256), nn.ReLU(), nn.Dropout(0.3),\n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, image, meta):\n        img_features = self.backbone(image)\n        img_features = self.pool(img_features).flatten(1)\n        meta_features = self.meta_fc(meta)\n        combined = torch.cat([img_features, meta_features], dim=1)\n        return self.classifier(combined)\n\n# --- Transforms ---\ndef get_tta_transforms():\n    \"\"\"Returns a list of safe and diverse TTA transforms.\"\"\"\n    base = [A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ToTensorV2()]\n    ttas = [\n        A.Compose(base),  # No augmentation\n        A.Compose([A.VerticalFlip(p=1.0)] + base),\n        A.Compose([A.Transpose(p=1.0)] + base),\n    ]\n    return ttas\n\nTRANSFORM_3D = A.Compose([A.Normalize(mean=0.5, std=0.5), ToTensorV2()])\nTRANSFORM_2D_INFERENCE = A.Compose([A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ToTensorV2()])\nTRANSFORM_2D_TTA = get_tta_transforms()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-08T11:32:22.25641Z","iopub.execute_input":"2025-09-08T11:32:22.256588Z","iopub.status.idle":"2025-09-08T11:32:22.276021Z","shell.execute_reply.started":"2025-09-08T11:32:22.256573Z","shell.execute_reply":"2025-09-08T11:32:22.275366Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Prediction and Model Loading","metadata":{}},{"cell_type":"code","source":"MODELS_3D = []\nMODELS_2D = {}\n\ndef load_all_models():\n    \"\"\"Loads all models and clears lists to ensure a clean state.\"\"\"\n    # Clear global lists to prevent duplication in interactive sessions\n    MODELS_3D.clear()\n    MODELS_2D.clear()\n    \n    # Pipeline 1: 3D Models\n    for fold in CFG.P1_MODEL_FOLDS:\n        model_path = os.path.join(CFG.P1_MODEL_DIR, f\"{CFG.P1_MODEL_NAME}_fold{fold}_best.pth\")\n        model = Timm3DModel(CFG.P1_MODEL_NAME)\n        sd = torch.load(model_path, map_location=device, weights_only=False)['model']\n        model.model.load_state_dict(sd)\n        model.to(device).eval()\n        MODELS_3D.append(model)\n    print(f\"Loaded {len(MODELS_3D)} 3D models.\")\n\n    # Pipeline 2: 2.5D Models\n    for name, path in CFG.P2_MODELS.items():\n        model_path = os.path.join(CFG.P2_MODEL_DIR, path)\n        model = MultiBackboneModel(name)\n        sd = torch.load(model_path, map_location=device, weights_only=False)['model_state_dict']\n        model.load_state_dict(sd)\n        model.to(device).eval()\n        MODELS_2D[name] = model\n    print(f\"Loaded {len(MODELS_2D)} 2.5D models.\")\n\ndef predict_3d_pipeline(volume_3d: np.ndarray) -> np.ndarray:\n    \"\"\"Runs the 3D pipeline with CPU-safe autocast.\"\"\"\n    # Autocast is now safely gated for CPU-only environments\n    with torch.no_grad(), autocast(enabled=(device.type == 'cuda')):\n        image_tensor = TRANSFORM_3D(image=volume_3d.transpose(1, 2, 0))['image']\n        image_tensor = image_tensor.unsqueeze(0).to(device)\n        all_preds = []\n        for model in MODELS_3D:\n            output = model(image_tensor)\n            all_preds.append(torch.sigmoid(output).cpu().numpy())\n    return np.mean(all_preds, axis=0).squeeze()\n\ndef predict_2d_pipeline(proj_image_2d: np.ndarray, metadata: Dict) -> np.ndarray:\n    \"\"\"Runs the 2.5D pipeline with real TTA and normalized weights.\"\"\"\n    with torch.no_grad(), autocast(enabled=(device.type == 'cuda')):\n        meta_tensor = torch.tensor([[metadata['age'] / 100.0, metadata['sex']]], dtype=torch.float32).to(device)\n        all_model_preds = []\n        for name, model in MODELS_2D.items():\n            tta_preds = []\n            if CFG.USE_TTA:\n                # Loop over the list of actual transforms\n                for t in TRANSFORM_2D_TTA:\n                    image_tensor = t(image=proj_image_2d)['image'].unsqueeze(0).to(device)\n                    output = model(image_tensor, meta_tensor)\n                    tta_preds.append(torch.sigmoid(output).cpu().numpy())\n            else:\n                image_tensor = TRANSFORM_2D_INFERENCE(image=proj_image_2d)['image'].unsqueeze(0).to(device)\n                output = model(image_tensor, meta_tensor)\n                tta_preds.append(torch.sigmoid(output).cpu().numpy())\n            \n            model_pred = np.mean(tta_preds, axis=0)\n            all_model_preds.append(model_pred)\n            \n        predictions = np.array(all_model_preds).squeeze(axis=1)\n        \n        # Safely get and normalize ensemble weights\n        weights = np.array([CFG.P2_ENSEMBLE_WEIGHTS.get(n, 0.0) for n in MODELS_2D.keys()], dtype=float)\n        if weights.sum() <= 0: # Fallback to equal weights\n            weights = np.ones_like(weights) / len(weights)\n        else:\n            weights /= weights.sum() # Normalize to sum to 1\n            \n        return np.average(predictions, weights=weights, axis=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-08T11:32:22.276742Z","iopub.execute_input":"2025-09-08T11:32:22.277226Z","iopub.status.idle":"2025-09-08T11:32:22.288521Z","shell.execute_reply.started":"2025-09-08T11:32:22.277208Z","shell.execute_reply":"2025-09-08T11:32:22.287965Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Main Execution","metadata":{}},{"cell_type":"code","source":"def predict(series_path: str) -> pl.DataFrame:\n    \"\"\"Top-level prediction function for the Kaggle server.\"\"\"\n    try:\n        # 1. Unified preprocessing\n        volume_3d, proj_image_2d, metadata = unified_preprocessor(series_path)\n\n        # 2. Run Pipeline 1 (3D)\n        preds_3d = predict_3d_pipeline(volume_3d)\n\n        # 3. Run Pipeline 2 (2.5D)\n        preds_2d = predict_2d_pipeline(proj_image_2d, metadata)\n\n        # 4. Final weighted ensemble of both pipelines\n        final_preds = (CFG.PIPELINE_3D_WEIGHT * preds_3d) + (CFG.PIPELINE_2D_WEIGHT * preds_2d)\n\n        # 5. Post-processing: ensure 'Aneurysm Present' is at least the max of others\n        max_location_prob = np.max(final_preds[:-1])\n        final_preds[-1] = np.max([final_preds[-1], max_location_prob])\n        \n        # Create output dataframe in the required format\n        return pl.DataFrame([final_preds.tolist()], schema=CFG.LABEL_COLS)\n\n    except Exception as e:\n        print(f\"Error processing {os.path.basename(series_path)}: {e}. Returning fallback.\")\n        return pl.DataFrame([[0.1] * len(CFG.LABEL_COLS)], schema=CFG.LABEL_COLS)\n    finally:\n        # Crucial memory and disk space cleanup\n        shared_dir = '/kaggle/shared'\n        shutil.rmtree(shared_dir, ignore_errors=True)\n        os.makedirs(shared_dir, exist_ok=True)\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        gc.collect()\n\n# Load all models at startup\nload_all_models()\n\n# Initialize and run the inference server\ninference_server = kaggle_evaluation.rsna_inference_server.RSNAInferenceServer(predict)\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    inference_server.serve()\nelse:\n    inference_server.run_local_gateway()\n    submission_df = pl.read_parquet('/kaggle/working/submission.parquet')\n    display(submission_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-08T11:32:22.289177Z","iopub.execute_input":"2025-09-08T11:32:22.289429Z","iopub.status.idle":"2025-09-08T11:32:59.708725Z","shell.execute_reply.started":"2025-09-08T11:32:22.289405Z","shell.execute_reply":"2025-09-08T11:32:59.708162Z"}},"outputs":[],"execution_count":null}]}