{"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"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":99552,"databundleVersionId":13190393,"sourceType":"competition"},{"sourceId":255258694,"sourceType":"kernelVersion"}],"dockerImageVersionId":31089,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **IA Detection: 8-Frame Image Inference Pipeline**","metadata":{}},{"cell_type":"markdown","source":"https://www.kaggle.com/code/stpeteishii/ia-detection-8-frame-image-training-pipeline\n\nhttps://www.kaggle.com/code/stpeteishii/ia-detection-8-frame-image-inference-pipeline","metadata":{}},{"cell_type":"markdown","source":"\n\n---\n\n# Introduction\n\nThis notebook implements an end-to-end inference pipeline for detecting intracranial aneurysms from multi-frame medical imaging data, specifically DICOM series. The pipeline leverages a deep learning model based on EfficientNet-B0, adapted to process multiple frames per study, to classify the presence of aneurysms across 14 distinct cerebral artery locations.\n\nKey features of this pipeline include:\n\n* **Multi-frame processing:** Uniformly samples and preprocesses 16 frames from each DICOM series, allowing the model to learn spatial and temporal patterns within volumetric medical scans.\n* **Domain-specific image windowing:** Applies intensity windowing based on imaging modality (e.g., CT, CTA, MRI) to enhance relevant tissue contrast.\n* **Robust DICOM handling:** Supports various photometric interpretations and pixel formats, ensuring reliable image extraction and normalization.\n* **EfficientNet-based architecture:** Utilizes a well-established convolutional backbone combined with temporal feature aggregation for accurate multi-label classification.\n* **Seamless integration with Kaggle inference environment:** Includes a server wrapper for efficient batch prediction and deployment in competition settings.\n\nThis notebook serves as the inference counterpart to the corresponding training pipeline, providing reliable predictions on unseen patient studies to support intracranial aneurysm detection efforts.\n\n---\n","metadata":{}},{"cell_type":"code","source":"# -------------------------------------------------------------\n# Required Imports\n# -------------------------------------------------------------\nimport os\nimport gc\nimport shutil\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport timm\nimport pydicom\nimport cv2\nimport polars as pl\nfrom torch.cuda.amp import autocast\nfrom pathlib import Path\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport kaggle_evaluation.rsna_inference_server\n\n# Device setup\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------\n# Hyper-parameters (must match training)\n# -------------------------------------------------------------\nclass CFG:\n    NUM_FRAMES   = 16\n    IMG_SIZE     = 224\n    NUM_CLASSES  = 14\n    MODEL_PATH   = \"/kaggle/input/ia-detection-8-frame-image-training-pipeline/efficientnet_b0_best.pth\"\n    USE_WINDOWING = True\n    USE_AMP       = True\n    BATCH_SIZE    = 1\n    WINDOW_PARAMS = {\n        'CT'   : (40, 80), 'CTA': (50,350),\n        'MRA': (600,1200), 'MRI' : (40,80),\n        'MRI T2' : (40,80), 'MRI T1post': (40,80)\n    }\n    LABELS = [\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    ]","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------\n# Model – identical to training (FIXED)\n# -------------------------------------------------------------\nclass MultiFrameEfficientNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = timm.create_model('efficientnet_b0',\n                                        pretrained=False,\n                                        num_classes=0,\n                                        global_pool='avg')\n        self.pool     = nn.AdaptiveAvgPool1d(1)\n        self.classifier = nn.Sequential(nn.Dropout(0.2),\n                                      nn.Linear(self.backbone.num_features,\n                                                CFG.NUM_CLASSES))\n\n    def forward(self, x):    # x: (B, T, C, H, W)\n        B, T, C, H, W = x.shape\n        feat = self.backbone(x.view(B*T,C,H,W))          # (B*T, F)\n        feat = feat.view(B,T,-1).transpose(1,2)          # (B, F, T)\n        agg  = self.pool(feat).squeeze(-1)               # (B, F)\n        return self.classifier(agg)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------\n# Helpers – DICOM → RGB-tensor (with memory optimizations)\n# -------------------------------------------------------------\ndef windowing(img, center, width):\n    mn, mx = center - width//2, center + width//2\n    img = np.clip(img, mn, mx)\n    return np.clip((img-mn)/(mx-mn+1e-7)*255,0,255).astype(np.uint8)\n\ndef get_window(mod):\n    return CFG.WINDOW_PARAMS.get(mod, CFG.WINDOW_PARAMS['CTA'])\n\ndef process_one(dcm_path: str, mod: str):\n    try:\n        # Only read necessary DICOM tags to reduce memory\n        ds = pydicom.dcmread(dcm_path, stop_before_pixels=True)\n        if 'PixelData' not in ds:  \n            return None\n        \n        # Now read pixel data separately\n        ds = pydicom.dcmread(dcm_path, force=True)\n        img = ds.pixel_array\n        del ds  # Free DICOM object immediately\n        \n        # Handle 3D images\n        if img.ndim == 3:\n            if hasattr(ds, 'PhotometricInterpretation') and ds.PhotometricInterpretation in (\"RGB\",\"YBR_FULL\"):\n                img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n            else: \n                img = img[:,:,0]\n        \n        # Handle MONOCHROME1 (inverted grayscale)\n        if hasattr(ds, 'PhotometricInterpretation') and ds.PhotometricInterpretation == 'MONOCHROME1':\n            img = 255 - img\n        \n        # Apply windowing or normalize\n        if CFG.USE_WINDOWING:\n            c, w = get_window(mod)\n            img = windowing(img, c, w)\n        else:\n            img_min = img.min()\n            img_max = img.max()\n            img = ((img-img_min)/(img_max-img_min+1e-7)*255).astype(np.uint8)\n        \n        # Resize and convert to RGB\n        img = cv2.resize(img, (CFG.IMG_SIZE, CFG.IMG_SIZE))\n        img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n        return img\n        \n    except Exception as e:\n        print(f\"❌ Error processing {dcm_path}: {e}\")\n        return None\n\ndef load_series(path: str):\n    try:\n        # Find all DICOM files with memory-efficient approach\n        files = []\n        for p in Path(path).rglob('*'):\n            if str(p).lower().endswith('.dcm'):\n                try:\n                    # Extract number for sorting without loading all files at once\n                    num = int(p.name.split('.')[-2]) if p.name.split('.')[-2].isdigit() else 0\n                    files.append((num, p))\n                except:\n                    files.append((0, p))\n        \n        if not files: \n            print(f\"❌ No DICOM files found in {path}\")\n            return torch.zeros(CFG.NUM_FRAMES, 3, CFG.IMG_SIZE, CFG.IMG_SIZE)\n        \n        # Sort by extracted number\n        files.sort(key=lambda x: x[0])\n        files = [p for (_, p) in files]\n        \n        # Get modality from first file (with limited loading)\n        try:\n            mod = pydicom.dcmread(files[0], stop_before_pixels=True).Modality\n        except:\n            mod = 'CTA'  # Default modality\n        \n        # Process images with memory cleanup\n        imgs = []\n        for p in files[:CFG.NUM_FRAMES*2]:  # Limit number of files processed\n            processed_img = process_one(str(p), mod)\n            if processed_img is not None:\n                imgs.append(processed_img)\n            if len(imgs) >= CFG.NUM_FRAMES*2:  # Don't need more than 2x frames\n                break\n        \n        if not imgs: \n            print(f\"❌ No valid images processed from {path}\")\n            return torch.zeros(CFG.NUM_FRAMES, 3, CFG.IMG_SIZE, CFG.IMG_SIZE)\n        \n        # Sample frames more efficiently\n        if len(imgs) >= CFG.NUM_FRAMES:\n            idxs = np.linspace(0, len(imgs)-1, CFG.NUM_FRAMES, dtype=int)\n        else:\n            repeat = CFG.NUM_FRAMES // len(imgs)\n            rem = CFG.NUM_FRAMES % len(imgs)\n            idxs = list(range(len(imgs))) * repeat + list(np.linspace(0, len(imgs)-1, rem, dtype=int))\n        \n        # Apply transformations with minimal memory usage\n        tfm = A.Compose([\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ])\n        \n        # Process images one by one to avoid memory spikes\n        result = []\n        for i in idxs:\n            transformed = tfm(image=imgs[i])['image']\n            result.append(transformed)\n            del transformed\n            gc.collect()\n        \n        # Clear processed images\n        del imgs\n        gc.collect()\n        \n        return torch.stack(result)\n        \n    except Exception as e:\n        print(f\"❌ Error loading series {path}: {e}\")\n        return torch.zeros(CFG.NUM_FRAMES, 3, CFG.IMG_SIZE, CFG.IMG_SIZE)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------\n# Model loading (with memory optimizations)\n# -------------------------------------------------------------\nMODEL = None\n\ndef get_model():\n    global MODEL\n    if MODEL: \n        return MODEL\n    \n    print(\"🔄 Loading model...\")\n    MODEL = MultiFrameEfficientNet().to(device)\n    MODEL.eval()\n    \n    try:\n        # Load checkpoint with minimal memory footprint\n        ckpt = torch.load(CFG.MODEL_PATH, map_location='cpu', weights_only=False)\n        \n        # Handle different checkpoint formats\n        state_dict = ckpt\n        if isinstance(ckpt, dict):\n            for key in ['model_state_dict', 'model', 'state_dict']:\n                if key in ckpt:\n                    state_dict = ckpt[key]\n                    break\n        \n        # Load state dict and clean up\n        MODEL.load_state_dict(state_dict, strict=False)\n        del ckpt, state_dict\n        gc.collect()\n        print(\"✅ Model loaded successfully\")\n        \n    except Exception as e:\n        print(f\"❌ Model loading failed: {e}\")\n        print(\"⚠️ Using randomly initialized weights\")\n    \n    # Warm-up inference with smaller batch\n    try:\n        with torch.no_grad():\n            # Use smaller dummy input for warm-up\n            dummy_input = torch.randn(1, min(4, CFG.NUM_FRAMES), 3, CFG.IMG_SIZE//2, CFG.IMG_SIZE//2).to(device)\n            \n            if hasattr(torch.amp, 'autocast'):\n                with torch.amp.autocast('cuda', enabled=CFG.USE_AMP):\n                    _ = MODEL(dummy_input)\n            else:\n                with autocast(enabled=CFG.USE_AMP):\n                    _ = MODEL(dummy_input)\n            \n            del dummy_input\n            torch.cuda.empty_cache()\n        print(\"✅ Model warm-up completed\")\n    except Exception as e:\n        print(f\"❌ Model warm-up failed: {e}\")\n    \n    return MODEL","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------\n# Prediction API (with fallback and memory/disk cleanup)\n# -------------------------------------------------------------\ndef predict(series_path: str) -> pl.DataFrame:\n    \"\"\"\n    Main prediction function for Kaggle RSNA Inference API.\n    Always returns a valid DataFrame with correct schema, even on error.\n    \"\"\"\n    try:\n        inp = load_series(series_path).unsqueeze(0).to(device)\n\n        with torch.no_grad():\n            if hasattr(torch.amp, 'autocast'):\n                with torch.amp.autocast('cuda', enabled=CFG.USE_AMP):\n                    logits = get_model()(inp)\n            else:\n                with autocast(enabled=CFG.USE_AMP):\n                    logits = get_model()(inp)\n\n        probs = torch.sigmoid(logits).cpu().numpy()[0]\n        probs = np.clip(probs, 0, 1)\n\n        print(f\"✅ Prediction successful for {os.path.basename(series_path)}\")\n\n    except Exception as e:\n        print(f\"❌ Prediction failed for {os.path.basename(series_path)}: {e}\")\n        probs = np.full(CFG.NUM_CLASSES, 0.3)\n\n    # Ensure no NaN/Inf\n    probs = np.nan_to_num(probs, nan=0.3, posinf=1.0, neginf=0.0)\n\n    # Create DataFrame with fixed schema & order\n    df = pl.DataFrame(\n        [probs.tolist()],\n        schema=CFG.LABELS,\n        orient=\"row\"\n    )\n\n    # Disk cleanup to avoid Kaggle error\n    shared_dir = '/kaggle/shared'\n    shutil.rmtree(shared_dir, ignore_errors=True)\n    os.makedirs(shared_dir, exist_ok=True)\n\n    # Memory cleanup\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n    gc.collect()\n\n    return df\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------------------------------------\n# Kaggle inference server wrapper (robust version)\n# -------------------------------------------------------------\ndef main():\n    print(\"=\"*70)\n    print(\"RSNA Intracranial Aneurysm Detection - Inference\")\n    print(\"=\"*70)\n    print(f\"Device: {device}\")\n    print(f\"Model: MultiFrameEfficientNet\")\n    print(f\"Frames: {CFG.NUM_FRAMES}\")\n    print(f\"Image size: {CFG.IMG_SIZE}x{CFG.IMG_SIZE}\")\n    print(f\"Classes: {CFG.NUM_CLASSES}\")\n    print(\"-\"*70)\n\n    try:\n        get_model()\n\n        server = kaggle_evaluation.rsna_inference_server.RSNAInferenceServer(predict)\n\n        if os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n            print(\"🔄 Running in competition mode...\")\n            server.serve()\n        else:\n            print(\"🔧 Running in local testing mode...\")\n            server.run_local_gateway()\n\n            sub_path = '/kaggle/working/submission.parquet'\n            if os.path.exists(sub_path):\n                try:\n                    df = pl.read_parquet(sub_path)\n                    print(f\"📄 Submission preview: {df.shape}\")\n                    print(df.head())\n                except Exception as e:\n                    print(f\"⚠️ Could not read submission file: {e}\")\n\n        print(\"=\"*70)\n        print(\"✅ Inference completed successfully!\")\n        print(\"=\"*70)\n\n    except Exception as e:\n        print(f\"💥 Critical error in main(): {e}\")\n        raise e\n\n    finally:\n        gc.collect()\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n\n        global MODEL\n        if MODEL is not None:\n            del MODEL\n            MODEL = None\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n---\n\n## Overall Pipeline Overview\n\n### 1. **Configuration (CFG class)**\n\n* Defines hyperparameters such as:\n\n  * Number of frames per series (`NUM_FRAMES = 16`)\n  * Input image size (`IMG_SIZE = 224x224`)\n  * Number of output classes (`NUM_CLASSES = 14`)\n  * Model checkpoint path\n  * Whether to apply windowing on images (for CT/MRI intensity normalization)\n  * Batch size and other relevant parameters\n* Specifies windowing parameters (center, width) per imaging modality (CT, CTA, MRI, etc.)\n* Defines label names for the 14 anatomical regions plus “Aneurysm Present”\n\n### 2. **Model Architecture (`MultiFrameEfficientNet`)**\n\n* Based on EfficientNet-B0 backbone (no pretrained weights used here)\n* Takes multi-frame input: tensor of shape (Batch, Time, Channels, Height, Width)\n* Processes each frame individually through EfficientNet to extract features\n* Aggregates features over time (frames) via adaptive average pooling\n* Final classifier outputs probabilities for 14 classes\n\n### 3. **DICOM Image Processing Helpers**\n\n* Functions to read DICOM files and convert them into RGB tensors ready for the model:\n\n  * Handles pixel intensity windowing (HU values clipping & normalization)\n  * Converts multi-channel images (RGB or YBR) to grayscale if needed\n  * Handles MONOCHROME1 images by inverting intensities\n  * Resizes images to the fixed size (224x224)\n  * Applies normalization and converts to PyTorch tensors\n\n* For each series folder, loads all DICOM files, sorts by slice index, samples or repeats frames to get exactly 16 frames per series.\n\n### 4. **Model Loading**\n\n* Loads model weights from checkpoint\n* Supports checkpoint dicts with or without `'model_state_dict'` key\n* Moves model to GPU if available and sets to eval mode\n* Runs a warm-up forward pass with dummy input to optimize performance\n\n### 5. **Prediction Function**\n\n* Takes a series path (folder of DICOM files)\n* Loads and preprocesses 16 frames into tensor batch\n* Runs inference with the model, applies sigmoid to get probabilities\n* Clips probabilities between 0 and 1\n* Returns predictions as a Polars DataFrame with labels as column names\n* Handles exceptions by returning default moderate probability scores if prediction fails\n\n### 6. **Kaggle Inference Server Wrapper**\n\n* Defines `main()` function that:\n\n  * Instantiates inference server with the `predict` function\n  * Runs server in either competition environment or local mode\n  * Performs cleanup of temporary files and clears GPU cache after execution\n\n---\n\nThis pipeline reads multi-frame DICOM series, preprocesses images with windowing and normalization, runs inference on an EfficientNet-based multi-frame model, and outputs probability predictions for intracranial aneurysm detection at 14 anatomical locations.\n","metadata":{}}]}