{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":99552,"databundleVersionId":13851420,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13334997,"sourceType":"datasetVersion","datasetId":8455177},{"sourceId":13350304,"sourceType":"datasetVersion","datasetId":8466874}],"dockerImageVersionId":31154,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n# ============================================================================\n# CELL 1: IMPORTS\n# ============================================================================\n\nimport os\nimport shutil\nfrom collections import defaultdict\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport pydicom\nfrom pathlib import Path\nfrom scipy.ndimage import zoom, rotate as scipy_rotate\nimport cv2\nfrom typing import List, Tuple, Dict\nimport gc\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport kaggle_evaluation.rsna_inference_server\n\nprint(\"✅ All imports successful\")\n\n# Device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"✅ Using device: {device}\")\nprint(f\"   Available GPUs: {torch.cuda.device_count()}\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 2: CONFIGURATION + POST-PROCESSING RULES\n# ============================================================================\n\n# Competition constants\nID_COL = 'SeriesInstanceUID'\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# Model configuration\nclass Config:\n    MODEL_NAME = 'tf_efficientnetv2_s.in21k_ft_in1k'\n    INPUT_SIZE = (32, 384, 384)\n    IN_CHANNELS = 32\n    NUM_CLASSES = 14\n    SIZE = 384\n    \n    # Ensemble settings\n    MODEL_DIR = '/kaggle/input/rsna-aneurysm-efficientnet-ensemble-5fold/ensemble_models'\n    N_FOLDS = 5\n    FOLDS = [0, 1, 2, 3, 4]\n    \n    # Post-processing settings\n    USE_POST_PROCESSING = True\n    \n    # Ensemble weights (equal or CV-based)\n    ENSEMBLE_WEIGHTS = None  # Equal weights\n\nconfig = Config()\n\n# ============================================================================\n# POST-PROCESSING FUNCTIONS\n# ============================================================================\n\ndef apply_consistency_rules(predictions: np.ndarray) -> np.ndarray:\n    \"\"\"\n    Apply logical consistency rules to predictions.\n    \n    Args:\n        predictions: (14,) array [13 location probs + 1 aneurysm prob]\n    \n    Returns:\n        processed predictions: (14,) array\n    \"\"\"\n    predictions = predictions.copy()\n    \n    location_probs = predictions[:-1]\n    aneurysm_prob = predictions[-1]\n    \n    # Rule 1: Aneurysm probability should be at least as high as max location\n    # Logic: If any location has high probability, aneurysm should too\n    max_location_prob = location_probs.max()\n    if aneurysm_prob < max_location_prob:\n        predictions[-1] = max(aneurysm_prob, max_location_prob * 0.95)\n    \n    # Rule 2: If all locations are very low, moderate aneurysm probability\n    # Logic: Can't have aneurysm if no location shows evidence\n    if max_location_prob < 0.15:\n        predictions[-1] = min(predictions[-1], 0.6)\n    \n    # Rule 3: If any location is very confident, boost aneurysm\n    # Logic: High location confidence implies aneurysm present\n    if max_location_prob > 0.75:\n        predictions[-1] = max(predictions[-1], 0.75)\n    \n    # Rule 4: Smooth extreme predictions slightly\n    # Logic: Avoid overconfident predictions\n    predictions = np.clip(predictions, 0.001, 0.999)\n    \n    # Rule 5: If multiple locations are moderately high, boost aneurysm\n    # Logic: Multiple affected arteries suggest aneurysm\n    high_locations = (location_probs > 0.4).sum()\n    if high_locations >= 3:\n        predictions[-1] = max(predictions[-1], 0.65)\n    \n    return predictions\n\ndef temperature_scaling(predictions: np.ndarray, temperature: float = 1.2) -> np.ndarray:\n    \"\"\"\n    Apply temperature scaling for calibration.\n    \n    Args:\n        predictions: probability predictions\n        temperature: scaling factor (>1 = less confident, <1 = more confident)\n    \n    Returns:\n        calibrated predictions\n    \"\"\"\n    epsilon = 1e-7\n    predictions = np.clip(predictions, epsilon, 1 - epsilon)\n    \n    # Convert to logits\n    logits = np.log(predictions / (1 - predictions))\n    \n    # Scale\n    scaled_logits = logits / temperature\n    \n    # Convert back to probabilities\n    scaled_probs = 1 / (1 + np.exp(-scaled_logits))\n    \n    return scaled_probs\n\ndef apply_post_processing(predictions: np.ndarray) -> np.ndarray:\n    \"\"\"\n    Apply all post-processing steps.\n    \n    Args:\n        predictions: (14,) array of raw predictions\n    \n    Returns:\n        processed predictions: (14,) array\n    \"\"\"\n    if not config.USE_POST_PROCESSING:\n        return predictions\n    \n    # Step 1: Consistency rules\n    predictions = apply_consistency_rules(predictions)\n    \n    # Step 2: Temperature scaling (slightly less confident)\n    predictions = temperature_scaling(predictions, temperature=1.15)\n    \n    # Step 3: Final clipping\n    predictions = np.clip(predictions, 0.001, 0.999)\n    \n    return predictions\n\nprint(\"=\" * 80)\nprint(\"✅ POST-PROCESSING CONFIGURATION\")\nprint(\"=\" * 80)\nprint(f\"Post-processing enabled: {config.USE_POST_PROCESSING}\")\nprint(f\"Rules:\")\nprint(f\"  1. Aneurysm prob ≥ max location prob (consistency)\")\nprint(f\"  2. Low locations → moderate aneurysm (logic)\")\nprint(f\"  3. High location → boost aneurysm (confidence)\")\nprint(f\"  4. Clip extremes (calibration)\")\nprint(f\"  5. Multiple locations → boost aneurysm (evidence)\")\nprint(f\"  6. Temperature scaling: 1.15 (slight smoothing)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 3: DICOM PREPROCESSOR\n# ============================================================================\n\nclass DICOMPreprocessor:\n    \"\"\"DICOM preprocessing for inference\"\"\"\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    def load_dicom_series(self, series_path: str) -> List:\n        \"\"\"Load DICOM files from series path\"\"\"\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\")\n        \n        return datasets\n    \n    def extract_slice_info(self, datasets: List) -> List[Dict]:\n        \"\"\"Extract slice position information\"\"\"\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            try:\n                ipp = np.array(getattr(ds, 'ImagePositionPatient', None))\n                iop = np.array(getattr(ds, 'ImageOrientationPatient', None))\n                n_vec = np.cross(iop[:3], iop[3:])\n                info['z_position'] = -float((ipp * n_vec).sum())\n            except:\n                info['z_position'] = float(i)\n            slice_info.append(info)\n        return slice_info\n    \n    def sort_slices(self, slice_info: List[Dict]) -> List[Dict]:\n        \"\"\"Sort slices by z-position\"\"\"\n        return sorted(slice_info, key=lambda x: x['z_position'])\n    \n    def apply_normalization(self, img: np.ndarray, modality: str) -> np.ndarray:\n        \"\"\"Apply normalization\"\"\"\n        p1, p99 = np.percentile(img, [1, 99])\n        \n        if modality == 'CT' or modality == 'CTA':\n            p1, p99 = 0, 500\n        \n        if p99 > p1:\n            normalized = np.clip(img, p1, p99)\n            normalized = (normalized - p1) / (p99 - p1)\n            return (normalized * 255).astype(np.uint8)\n        else:\n            img_min, img_max = img.min(), img.max()\n            if img_max > img_min:\n                normalized = (img - img_min) / (img_max - img_min)\n                return (normalized * 255).astype(np.uint8)\n            else:\n                return np.zeros_like(img, dtype=np.uint8)\n    \n    def extract_pixel_array(self, ds) -> np.ndarray:\n        \"\"\"Extract pixel array from DICOM\"\"\"\n        img = ds.pixel_array.astype(np.float32)\n        if img.ndim == 3:\n            frame_idx = img.shape[0] // 2\n            img = img[frame_idx]\n        if img.ndim == 3 and img.shape[-1] == 3:\n            img = cv2.cvtColor(img.astype(np.uint8), cv2.COLOR_RGB2GRAY).astype(np.float32)\n        return img\n    \n    def resize_volume_3d(self, volume: np.ndarray) -> np.ndarray:\n        \"\"\"Resize 3D volume to target shape\"\"\"\n        current_shape = volume.shape\n        target_shape = (self.target_depth, self.target_height, self.target_width)\n        \n        if current_shape == target_shape:\n            return volume\n        \n        zoom_factors = [target_shape[i] / current_shape[i] for i in range(3)]\n        resized = zoom(volume, zoom_factors, order=1, mode='nearest')\n        resized = resized[:self.target_depth, :self.target_height, :self.target_width]\n        \n        pad_width = [\n            (0, max(0, self.target_depth - resized.shape[0])),\n            (0, max(0, self.target_height - resized.shape[1])),\n            (0, max(0, self.target_width - resized.shape[2]))\n        ]\n        \n        if any(pw[1] > 0 for pw in pad_width):\n            resized = np.pad(resized, pad_width, mode='edge')\n        \n        return resized.astype(np.uint8)\n    \n    def process_series(self, series_path: str) -> np.ndarray:\n        \"\"\"Main processing pipeline\"\"\"\n        datasets = self.load_dicom_series(series_path)\n        \n        first_ds = datasets[0]\n        first_img = first_ds.pixel_array\n        \n        if len(datasets) == 1 and first_img.ndim == 3:\n            return self._process_3d_dicom(first_ds)\n        else:\n            return self._process_2d_dicoms(datasets)\n    \n    def _process_3d_dicom(self, ds) -> np.ndarray:\n        \"\"\"Process single 3D DICOM\"\"\"\n        volume = ds.pixel_array.astype(np.float32)\n        modality = getattr(ds, 'Modality', 'CT')\n        \n        processed_slices = []\n        for i in range(volume.shape[0]):\n            slice_img = volume[i]\n            processed = self.apply_normalization(slice_img, modality)\n            processed_slices.append(processed)\n        \n        volume = np.stack(processed_slices, axis=0)\n        return self.resize_volume_3d(volume)\n    \n    def _process_2d_dicoms(self, datasets: List) -> np.ndarray:\n        \"\"\"Process multiple 2D DICOMs\"\"\"\n        slice_info = self.extract_slice_info(datasets)\n        sorted_slices = self.sort_slices(slice_info)\n        \n        first_ds = sorted_slices[0]['dataset']\n        modality = getattr(first_ds, 'Modality', 'CT')\n        \n        processed_slices = []\n        for slice_data in sorted_slices:\n            ds = slice_data['dataset']\n            img = self.extract_pixel_array(ds)\n            processed = self.apply_normalization(img, modality)\n            resized = cv2.resize(processed, (self.target_width, self.target_height))\n            processed_slices.append(resized)\n        \n        volume = np.stack(processed_slices, axis=0)\n        return self.resize_volume_3d(volume)\n\npreprocessor = DICOMPreprocessor(target_shape=config.INPUT_SIZE)\nprint(\"✅ DICOM Preprocessor ready\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 4: LOAD MODELS\n# ============================================================================\n\nprint(\"Loading ensemble models...\")\n\n# Storage for models\nMODELS = {}\n\n# Transforms (base)\ntransform = A.Compose([\n    A.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n        max_pixel_value=255.0\n    ),\n    ToTensorV2(),\n])\n\n# Load all fold models\nfor fold in config.FOLDS:\n    model_path = os.path.join(config.MODEL_DIR, f'fold{fold}_best_model.pth')\n    \n    print(f\"Loading fold {fold}...\")\n    \n    # Load checkpoint\n    checkpoint = torch.load(model_path, map_location='cpu', weights_only=False)\n    \n    # Create model\n    model = timm.create_model(\n        config.MODEL_NAME,\n        pretrained=False,\n        num_classes=config.NUM_CLASSES,\n        in_chans=config.IN_CHANNELS\n    )\n    \n    # Load weights\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model = model.to(device)\n    model.eval()\n    \n    MODELS[fold] = model\n    print(f\"  ✅ Fold {fold} loaded (AUC: {checkpoint['metric']:.4f})\")\n\nprint(f\"\\n✅ All {len(MODELS)} models loaded!\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 5: PREDICT FUNCTION WITH POST-PROCESSING\n# ============================================================================\n\ndef predict_single_model(model: nn.Module, image: np.ndarray) -> np.ndarray:\n    \"\"\"Make prediction with a single model\"\"\"\n    # Transpose for transforms: (D,H,W) -> (H,W,D)\n    image = image.transpose(1, 2, 0)\n    \n    # Apply transforms\n    transformed = transform(image=image)\n    image_tensor = transformed['image']  # (32, 384, 384)\n    image_tensor = image_tensor.unsqueeze(0).to(device)  # (1, 32, 384, 384)\n    \n    with torch.no_grad():\n        output = model(image_tensor)\n        return torch.sigmoid(output).cpu().numpy().squeeze()\n\ndef predict_ensemble(image: np.ndarray) -> np.ndarray:\n    \"\"\"\n    Make ensemble prediction.\n    \n    Args:\n        image: (D, H, W) numpy array (preprocessed volume)\n    \n    Returns:\n        final_prediction: (14,) numpy array\n    \"\"\"\n    all_fold_predictions = []\n    \n    # Get predictions from all folds\n    for fold, model in MODELS.items():\n        pred = predict_single_model(model, image)\n        all_fold_predictions.append(pred)\n    \n    # Equal-weight average across folds\n    predictions = np.array(all_fold_predictions)\n    ensemble_pred = np.mean(predictions, axis=0)\n    \n    # Apply post-processing\n    final_pred = apply_post_processing(ensemble_pred)\n    \n    return final_pred\n\ndef predict(series_path: str) -> pl.DataFrame:\n    \"\"\"\n    Main prediction function for API.\n    \"\"\"\n    series_id = os.path.basename(series_path)\n    \n    try:\n        # Process DICOM series\n        volume = preprocessor.process_series(series_path)\n        \n        # Make ensemble prediction with post-processing\n        final_pred = predict_ensemble(volume)\n        \n        # Create output dataframe\n        result = pl.DataFrame(\n            data=[[series_id] + final_pred.tolist()],\n            schema=[ID_COL] + LABEL_COLS,\n            orient='row',\n        )\n        \n        return result.drop(ID_COL)\n        \n    except Exception as e:\n        # Conservative fallback\n        conservative_preds = [0.1] * len(LABEL_COLS)\n        result = pl.DataFrame(\n            data=[conservative_preds],\n            schema=LABEL_COLS,\n            orient='row',\n        )\n        return result\n    finally:\n        # Cleanup\n        shutil.rmtree('/kaggle/shared', ignore_errors=True)\n        os.makedirs('/kaggle/shared', exist_ok=True)\n        if torch.cuda.is_available():\n            torch.cuda.empty_cache()\n        gc.collect()\n\nprint(\"✅ Predict function defined (WITH POST-PROCESSING)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CELL 6: START INFERENCE SERVER\n# ============================================================================\n\n# Create inference server\ninference_server = kaggle_evaluation.rsna_inference_server.RSNAInferenceServer(predict)\n\n# Run server\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    print(\"🏆 COMPETITION MODE: Starting inference server...\")\n    inference_server.serve()\nelse:\n    print(\"🧪 LOCAL TEST MODE: Running on sample data...\")\n    inference_server.run_local_gateway()\n    \n    print(\"\\n📊 Sample predictions:\")\n    display(pl.read_parquet('/kaggle/working/submission.parquet'))\n    print(\"\\n✅ Local test complete! (WITH POST-PROCESSING)\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}