{"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":12780021,"sourceType":"datasetVersion","datasetId":8079690},{"sourceId":14045202,"sourceType":"datasetVersion","datasetId":8941532},{"sourceId":14048010,"sourceType":"datasetVersion","datasetId":8943298}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport shutil\nfrom collections import defaultdict\nfrom pathlib import Path\nimport gc\n\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nimport pydicom\nimport cv2\nfrom scipy import ndimage\nfrom typing import List, Dict, Tuple\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\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"7ec6cc6e-3932-4c68-984d-9d6ed96863cc","_cell_guid":"065d6250-462b-461f-8297-39f8edfcf0f5","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# CONFIGURATION\n# ====================================================\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\nclass Config:\n    model_dir = '/kaggle/input/rsna-aneurysm-5fold-models/models'\n    cache_dir = '/kaggle/input/rsna-aneurysm-cache/cache'\n    \n    model_name = \"tf_efficientnetv2_s.in21k_ft_in1k\"\n    size = 384\n    in_chans = 32\n    num_classes = 14\n    target_shape = (32, 384, 384)\n    folds = [0, 1, 2, 3, 4]\n    device = 'cuda' if torch.cuda.is_available() else 'cpu'\n    use_cache = True\n\nCFG = Config()","metadata":{"_uuid":"a72b9ab0-98dd-48b9-9a9d-a04162846fdb","_cell_guid":"7dd894e1-61e3-45b1-85ff-1f6d55ee347c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# DICOM PREPROCESSING\n# ====================================================\nclass DICOMPreprocessorKaggle:\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[pydicom.Dataset]:\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\")\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[pydicom.Dataset]) -> List[Dict]:\n        slice_info = []\n        \n        for i, ds in enumerate(datasets):\n            info = {\n                'dataset': ds,\n                'index': i,\n                'instance_number': getattr(ds, 'InstanceNumber', i),\n            }\n            \n            try:\n                position = getattr(ds, 'ImagePositionPatient', None)\n                if position is not None and len(position) >= 3:\n                    info['z_position'] = float(position[2])\n                elif hasattr(ds, \"SliceLocation\"):\n                    info['z_position'] = float(getattr(ds, \"SliceLocation\"))\n                else:\n                    info['z_position'] = float(info['instance_number'])\n            except:\n                info['z_position'] = float(i)\n            \n            slice_info.append(info)\n        \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    def apply_windowing_or_normalize(self, img: np.ndarray) -> np.ndarray:\n        p1, p99 = np.percentile(img, [1, 99])\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: pydicom.Dataset) -> np.ndarray:\n        img = ds.pixel_array.astype(np.float32)\n        \n        if img.ndim == 3:\n            frame_idx = img.shape[0] // 2\n            img = img[frame_idx]\n        \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        \n        slope = float(getattr(ds, 'RescaleSlope', 1))\n        intercept = float(getattr(ds, 'RescaleIntercept', 0))\n        img = img * slope + intercept\n        \n        return img\n    \n    def resize_volume_3d(self, volume: np.ndarray) -> np.ndarray:\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_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        \n        if any(pw[1] > 0 for pw in pad_width):\n            resized_volume = np.pad(resized_volume, pad_width, mode='edge')\n        \n        return resized_volume.astype(np.uint8)\n    \n    def process_series(self, series_path: str) -> np.ndarray:\n        datasets = 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            return self._process_single_3d_dicom(first_ds)\n        else:\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        \n        slope = float(getattr(ds, 'RescaleSlope', 1))\n        intercept = float(getattr(ds, 'RescaleIntercept', 0))\n        volume = volume * slope + intercept\n        \n        processed_slices = []\n        for i in range(volume.shape[0]):\n            processed_img = self.apply_windowing_or_normalize(volume[i])\n            processed_slices.append(processed_img)\n        \n        volume = np.stack(processed_slices, 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        \n        processed_slices = []\n        for slice_data in sorted_slices:\n            ds = slice_data['dataset']\n            img = self.extract_pixel_array(ds)\n            processed_img = self.apply_windowing_or_normalize(img)\n            resized_img = cv2.resize(processed_img, (self.target_width, self.target_height))\n            processed_slices.append(resized_img)\n\n        volume = np.stack(processed_slices, axis=0)\n        return self.resize_volume_3d(volume)\n\ndef process_dicom_series_safe(series_path: str, target_shape: Tuple[int, int, int] = (32, 384, 384)) -> np.ndarray:\n    try:\n        preprocessor = DICOMPreprocessorKaggle(target_shape=target_shape)\n        return preprocessor.process_series(series_path)\n    finally:\n        gc.collect()","metadata":{"_uuid":"6bb5d611-bcf9-4db4-b8bf-a5eb4b7f25c9","_cell_guid":"72d8d2cd-3e77-49a5-b27a-7c326b653ae2","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ====================================================\n# LOAD MODELS \n# ====================================================\ndef load_model(fold):\n    model = timm.create_model(\n        CFG.model_name,\n        pretrained=False,\n        num_classes=CFG.num_classes,\n        in_chans=CFG.in_chans\n    )\n    \n    model_path = f'{CFG.model_dir}/{CFG.model_name}_fold{fold}_best.pth'\n    checkpoint = torch.load(model_path, map_location='cpu', weights_only=False)  # ← FIXED\n    model.load_state_dict(checkpoint['model'])\n    model.to(CFG.device)\n    model.eval()\n    \n    return model\n\nprint('Loading models...')\nmodels = [load_model(fold) for fold in CFG.folds]\nprint(f'Loaded {len(models)} models')\n\n# Transform\ntransform = A.Compose([\n    A.Resize(CFG.size, CFG.size),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ToTensorV2(),\n])\n\n# ====================================================\n# PREDICT FUNCTION (CALLED PER SERIES)\n# ====================================================\ndef predict(series_path: str) -> pl.DataFrame | pd.DataFrame:\n    \"\"\"Make a prediction for one series.\"\"\"\n    \n    series_id = os.path.basename(series_path)\n    \n    try:\n        # Check cache first\n        cache_path = Path(CFG.cache_dir) / f\"{series_id}.npy\"\n        \n        if CFG.use_cache and cache_path.exists():\n            volume = np.load(cache_path)\n        else:\n            volume = process_dicom_series_safe(series_path, CFG.target_shape)\n        \n        # Preprocess\n        volume = volume.transpose(1, 2, 0)\n        image = transform(image=volume)['image']\n        image = image.unsqueeze(0).to(CFG.device)\n        \n        # Ensemble prediction\n        with torch.no_grad(), autocast():\n            all_preds = []\n            for model in models:\n                output = model(image)\n                pred = torch.sigmoid(output).cpu().numpy()[0]\n                all_preds.append(pred)\n        \n        final_pred = np.mean(all_preds, axis=0)\n        \n        # Create result dataframe\n        predictions = pl.DataFrame(\n            data=[[series_id] + final_pred.tolist()],\n            schema=[ID_COL, *LABEL_COLS],\n            orient='row',\n        )\n        \n    except Exception as e:\n        print(f'Error processing {series_id}: {e}')\n        # Return default predictions\n        predictions = pl.DataFrame(\n            data=[[series_id] + [0.5] * len(LABEL_COLS)],\n            schema=[ID_COL, *LABEL_COLS],\n            orient='row',\n        )\n    \n    finally:\n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    # Verify format\n    if isinstance(predictions, pl.DataFrame):\n        assert predictions.columns == [ID_COL, *LABEL_COLS]\n    elif isinstance(predictions, pd.DataFrame):\n        assert (predictions.columns == [ID_COL, *LABEL_COLS]).all()\n    else:\n        raise TypeError('The predict function must return a DataFrame')\n\n    # IMPORTANT: Prevent disk space errors\n    shutil.rmtree('/kaggle/shared', ignore_errors=True)\n    \n    # Return WITHOUT the ID column (as required)\n    return predictions.drop(ID_COL)\n\n# ====================================================\n# START INFERENCE SERVER\n# ====================================================\ninference_server = kaggle_evaluation.rsna_inference_server.RSNAInferenceServer(predict)\n\nif os.getenv('KAGGLE_IS_COMPETITION_RERUN'):\n    inference_server.serve()\nelse:\n    inference_server.run_local_gateway()\n    display(pl.read_parquet('/kaggle/working/submission.parquet'))","metadata":{"_uuid":"6f3860b8-34a5-4846-8539-8d87cab1152c","_cell_guid":"bb4a1d3a-b79a-4afa-8b97-4b3defe03525","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}},"outputs":[],"execution_count":null}]}