{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","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":13851420,"sourceType":"competition"},{"sourceId":14824555,"sourceType":"datasetVersion","datasetId":9480816}],"dockerImageVersionId":31259,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ============================================================\n# RSNA Intracranial Aneurysm Detection - INFERENCE ONLY NOTEBOOK\n# Kaggle Submission Version (Fixed & Optimized)\n# Uses trained fold models stored in Kaggle Dataset: rsma2025\n# ============================================================\n\nimport os\nimport numpy as np\nimport cv2\nimport pydicom\nfrom scipy import ndimage\nfrom pathlib import Path\nimport torch\nimport torch.nn as nn\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport polars as pl\nimport kaggle_evaluation.rsna_inference_server\n\n\n# =========================\n# Device\n# =========================\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)\n\n\n# =========================\n# Constants\n# =========================\nIMAGE_SIZE = 384  # Match training config\nNUM_SLICES = 32\nNUM_FOLDS = 3\n\n# ✅ Correct Kaggle input path (your dataset)\nMODEL_DIR = \"/kaggle/input/datasets/amit393/rsma2025\"\n\n\n# =========================\n# Label Columns\n# =========================\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# =========================\n# Validation Transform\n# =========================\nVAL_TRANSFORM = A.Compose([\n    A.Resize(IMAGE_SIZE, IMAGE_SIZE),\n    A.Normalize(mean=[0.5], std=[0.5]),\n    ToTensorV2()\n])\n\n\n# =========================\n# DICOM Preprocessing (FIXED VERSION)\n# Matches the exact preprocessing from training pipeline\n# =========================\nclass Preprocessing:\n    \"\"\"Fixed preprocessing class matching training pipeline.\"\"\"\n    \n    def __init__(self, target_shape=(NUM_SLICES, IMAGE_SIZE, IMAGE_SIZE)):\n        self.target_depth, self.target_height, self.target_width = target_shape\n\n    def load_dicom_series(self, path: str):\n        \"\"\"Load DICOM files with deferred pixel data loading.\"\"\"\n        series_path = Path(path)\n        dicom_files = sorted(series_path.glob(\"*.dcm\"))\n        if len(dicom_files) == 0:\n            raise ValueError(f\"No DICOM files in {series_path}\")\n        # FIX: defer_size reduces memory usage\n        return [pydicom.dcmread(str(f), force=True, defer_size=256) for f in dicom_files]\n\n    def extract_slice_info(self, datasets):\n        \"\"\"Extract z-position metadata.\"\"\"\n        slice_info = []\n        for i, ds in enumerate(datasets):\n            z = getattr(ds, \"ImagePositionPatient\", [0, 0, i])[2]\n            slice_info.append({'dataset': ds, 'z_position': float(z)})\n        return slice_info\n\n    def sort_slices(self, slice_info):\n        \"\"\"Sort slices by z-position.\"\"\"\n        return sorted(slice_info, key=lambda x: x['z_position'])\n\n    def apply_windowing_or_normalize(self, img: np.ndarray) -> np.ndarray:\n        \"\"\"Percentile windowing then normalize to [0, 255].\"\"\"\n        p1, p99 = np.percentile(img, [1, 99])\n        img = np.clip(img, p1, p99)\n        img = (img - p1) / (p99 - p1 + 1e-6)\n        return (img * 255).astype(np.uint8)\n\n    def extract_pixel_array(self, ds: pydicom.Dataset) -> np.ndarray:\n        \"\"\"\n        Extract pixel array with proper HU conversion.\n        FIX: Keep float32 throughout, only convert RGB->gray in float space.\n        \"\"\"\n        img = ds.pixel_array.astype(np.float32)\n        slope = float(getattr(ds, 'RescaleSlope', 1))\n        intercept = float(getattr(ds, 'RescaleIntercept', 0))\n        img = img * slope + intercept  # HU values\n\n        # FIX: Convert RGB to grayscale in float space, not uint8\n        if img.ndim == 3 and img.shape[-1] == 3:\n            img_u8 = np.clip(img, 0, 255).astype(np.uint8)\n            img = cv2.cvtColor(img_u8, cv2.COLOR_RGB2GRAY).astype(np.float32)\n\n        # FIX: Handle multi-frame DICOMs (take first frame)\n        if img.ndim == 3:\n            img = img[0]\n\n        return img  # Always 2D (H, W)\n\n    def resize_volume_3d(self, volume: np.ndarray) -> np.ndarray:\n        \"\"\"\n        Resize 3D volume to target shape in one pass.\n        FIX: Single ndimage.zoom call instead of per-slice resize.\n        \"\"\"\n        if volume.ndim != 3:\n            raise ValueError(f\"Expected 3D volume, got shape {volume.shape}\")\n        \n        zoom = [t / c for t, c in zip(\n            (self.target_depth, self.target_height, self.target_width),\n            volume.shape\n        )]\n        volume = ndimage.zoom(volume, zoom, order=1)\n        return volume.astype(np.uint8)\n\n    def process_series(self, series_path: str) -> np.ndarray:\n        \"\"\"\n        Main preprocessing pipeline.\n        Returns: (NUM_SLICES, IMAGE_SIZE, IMAGE_SIZE) volume.\n        \"\"\"\n        datasets = self.load_dicom_series(series_path)\n        slices = self.sort_slices(self.extract_slice_info(datasets))\n        \n        volume = []\n        for s in slices:\n            img = self.extract_pixel_array(s['dataset'])\n            # FIX: Skip non-2D slices\n            if img.ndim != 2:\n                continue\n            img = self.apply_windowing_or_normalize(img)\n            volume.append(img)\n\n        if len(volume) == 0:\n            raise ValueError(\"No valid 2D slices found\")\n\n        volume = np.stack(volume, axis=0)  # (D_orig, H_orig, W_orig)\n        volume = self.resize_volume_3d(volume)  # → (32, 384, 384)\n        return volume\n\n\n# Global preprocessor instance\nPREPROCESSOR = Preprocessing()\n\n\ndef load_and_preprocess_dicom_series(series_dir: str) -> np.ndarray:\n    \"\"\"\n    Wrapper function for DICOM preprocessing.\n    Returns middle slice as (1, H, W) for model input.\n    \"\"\"\n    try:\n        volume = PREPROCESSOR.process_series(series_dir)\n        \n        # FIX: Use middle slice matching training pipeline\n        mid = volume.shape[0] // 2\n        middle_slice = volume[mid]  # (H, W)\n        \n        return middle_slice[None, :, :]  # (1, H, W)\n        \n    except Exception as e:\n        print(f\"[WARN] Preprocessing failed for {series_dir}: {e}\")\n        # Return zero image on failure\n        return np.zeros((1, IMAGE_SIZE, IMAGE_SIZE), dtype=np.uint8)\n\n\n# =========================\n# Model Definition (FIXED)\n# Matches training architecture exactly\n# =========================\nclass EfficientNetRSNA(nn.Module):\n    \"\"\"Fixed model architecture matching training.\"\"\"\n    \n    def __init__(self, num_classes: int = len(LABEL_COLS)):\n        super().__init__()\n        \n        # Backbone with in_chans=1 for grayscale\n        self.backbone = timm.create_model(\n            \"tf_efficientnetv2_s\",\n            pretrained=False,\n            in_chans=1,\n            num_classes=0\n        )\n        \n        # FIX: Match exact head architecture from training\n        self.head = nn.Sequential(\n            nn.Linear(self.backbone.num_features, 512),\n            nn.BatchNorm1d(512),\n            nn.GELU(),\n            nn.Dropout(0.4),\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.GELU(),\n            nn.Dropout(0.4),\n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        features = self.backbone(x)\n        return self.head(features)\n\n\n# =========================\n# Load Ensemble Models\n# =========================\nENSEMBLE_MODELS = None\n\ndef load_ensemble_models():\n    \"\"\"Load all fold models into memory.\"\"\"\n    models = []\n\n    print(\"\\nChecking MODEL_DIR:\", MODEL_DIR)\n    if os.path.exists(MODEL_DIR):\n        print(\"Files inside MODEL_DIR:\", os.listdir(MODEL_DIR))\n    else:\n        raise FileNotFoundError(f\"MODEL_DIR not found: {MODEL_DIR}\")\n\n    for fold_index in range(NUM_FOLDS):\n        model_path = os.path.join(MODEL_DIR, f\"best_model_fold{fold_index}.pth\")\n\n        if not os.path.exists(model_path):\n            raise FileNotFoundError(f\"❌ Missing model file: {model_path}\")\n\n        # FIX: Use correct model class\n        model = EfficientNetRSNA().to(DEVICE)\n        \n        # FIX: Add weights_only=True for security\n        state_dict = torch.load(model_path, map_location=DEVICE, weights_only=True)\n        model.load_state_dict(state_dict)\n        model.eval()\n        models.append(model)\n\n        print(f\"✅ Loaded model fold {fold_index}: {model_path}\")\n\n    print(f\"\\n✅ Loaded {len(models)} models successfully.\")\n    return models\n\n\n# =========================\n# Kaggle Predict Function\n# =========================\ndef predict(series_path: str) -> pl.DataFrame:\n    \"\"\"\n    Main prediction function called by Kaggle inference server.\n    \n    Args:\n        series_path: Path to directory containing DICOM series\n        \n    Returns:\n        Polars DataFrame with predictions for all 14 labels\n    \"\"\"\n    global ENSEMBLE_MODELS\n\n    # Lazy load models on first call\n    if ENSEMBLE_MODELS is None:\n        ENSEMBLE_MODELS = load_ensemble_models()\n\n    # Preprocess DICOM series → (1, H, W)\n    image_np = load_and_preprocess_dicom_series(series_path)\n\n    # FIX: Proper shape handling for albumentations\n    # (1, H, W) → (H, W, 1) for transform → (1, H, W) tensor\n    image_for_transform = image_np.transpose(1, 2, 0)  # (H, W, 1)\n    image_tensor = VAL_TRANSFORM(image=image_for_transform)[\"image\"]\n    image_tensor = image_tensor.unsqueeze(0).to(DEVICE)  # (1, 1, H, W)\n\n    # Ensemble prediction across all folds\n    ensemble_probs = []\n\n    with torch.no_grad():\n        # FIX: Use AMP if on GPU\n        if DEVICE.type == \"cuda\":\n            with torch.amp.autocast(\"cuda\"):\n                for model in ENSEMBLE_MODELS:\n                    logits = model(image_tensor)\n                    probs = torch.sigmoid(logits).cpu().numpy().flatten()\n                    ensemble_probs.append(probs)\n        else:\n            for model in ENSEMBLE_MODELS:\n                logits = model(image_tensor)\n                probs = torch.sigmoid(logits).cpu().numpy().flatten()\n                ensemble_probs.append(probs)\n\n    # Average predictions across folds\n    final_probs = np.mean(ensemble_probs, axis=0).astype(np.float64)\n    \n    # FIX: Clip to valid probability range\n    final_probs = np.clip(final_probs, 0.0, 1.0)\n\n    # Create output DataFrame\n    output_data = {col: [float(val)] for col, val in zip(LABEL_COLS, final_probs)}\n    return pl.DataFrame(output_data)\n\n\n# =========================\n# Start Kaggle Inference Server\n# =========================\nprint(\"\\n🚀 Starting Kaggle RSNA Inference Server...\")\n\ninference_server = kaggle_evaluation.rsna_inference_server.RSNAInferenceServer(predict)\n\nif os.getenv(\"KAGGLE_IS_COMPETITION_RERUN\"):\n    print(\"✅ Running in OFFICIAL rerun mode...\")\n    inference_server.serve()\nelse:\n    print(\"🧪 Running in LOCAL test mode...\")\n    inference_server.run_local_gateway()\n\nprint(\"✅ Done.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}