{"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"},"jupytext":{"cell_metadata_filter":"-all","main_language":"python","notebook_metadata_filter":"-all"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":99552,"databundleVersionId":13441085}],"dockerImageVersionId":31089,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom typing import List, Tuple, Optional, Union\n\nImageLike = Union[str, np.ndarray, Image.Image]\n\ndef plot_image_grid(\n    images: List[ImageLike],\n    grid_size: Tuple[int, int],\n    figsize: Optional[Tuple[float, float]] = None,\n    titles: Optional[List[str]] = None,\n    cmap: Optional[str] = None,\n    pad_value: Union[int, float] = 0,\n    resize_to: Optional[Tuple[int, int]] = None,\n    interpolation: str = \"nearest\",\n    show: bool = True,\n    save_path: Optional[str] = None,\n) -> None:\n    nrows, ncols = grid_size\n    n_cells = nrows * ncols\n    if len(images) == 0:\n        raise ValueError(\"`images` list is empty.\")\n\n    # Helper: load/convert one image to numpy array\n    def _to_numpy(img: ImageLike) -> np.ndarray:\n        if isinstance(img, str):\n            pil = Image.open(img)\n            pil = pil.convert(\"RGBA\") if pil.mode == \"P\" else pil\n        elif isinstance(img, Image.Image):\n            pil = img\n        elif isinstance(img, np.ndarray):\n            arr = img\n            # If it's already a numpy array, return copy to avoid side effects\n            return np.asarray(arr)\n        else:\n            raise TypeError(f\"Unsupported image type: {type(img)}\")\n\n        if resize_to is not None:\n            pil = pil.resize(resize_to[::-1], Image.BILINEAR)  # PIL expects (width, height)\n        return np.asarray(pil)\n\n    # Convert provided images\n    imgs = []\n    for im in images[:n_cells]:\n        imgs.append(_to_numpy(im))\n\n    # Determine base shape for padding/resizing if necessary\n    if len(imgs) == 0:\n        raise ValueError(\"No convertible images found.\")\n    base_h, base_w = imgs[0].shape[:2]\n\n    # If resize_to is set we already resized. If not, we will pad images to base shape.\n    if resize_to is None:\n        for i, im in enumerate(imgs):\n            h, w = im.shape[:2]\n            if (h, w) != (base_h, base_w):\n                # pad to base size (centered)\n                if im.ndim == 2:\n                    padded = np.full((base_h, base_w), pad_value, dtype=im.dtype)\n                else:\n                    ch = im.shape[2]\n                    padded = np.full((base_h, base_w, ch), pad_value, dtype=im.dtype)\n                # compute offsets\n                off_h = (base_h - h) // 2\n                off_w = (base_w - w) // 2\n                padded[off_h:off_h + h, off_w:off_w + w, ...] = im\n                imgs[i] = padded\n\n    # Pad list of images if fewer than grid cells\n    if len(imgs) < n_cells:\n        # create blank image matching shape of first image\n        first = imgs[0]\n        if first.ndim == 2:\n            blank = np.full((base_h, base_w), pad_value, dtype=first.dtype)\n        else:\n            blank = np.full((base_h, base_w, first.shape[2]), pad_value, dtype=first.dtype)\n        imgs += [blank] * (n_cells - len(imgs))\n\n    # Prepare figure\n    if figsize is None:\n        figsize = (ncols * 2, nrows * 2)\n    fig, axes = plt.subplots(nrows, ncols, figsize=figsize)\n    # Flatten axes for simple indexing (works for 1xN, Nx1, NxM)\n    if isinstance(axes, np.ndarray):\n        axes_flat = axes.flatten()\n    else:\n        axes_flat = [axes]\n\n    for idx, ax in enumerate(axes_flat[:n_cells]):\n        img = imgs[idx]\n        # If grayscale (2D), use cmap (or 'gray' default)\n        if img.ndim == 2:\n            ax.imshow(img, cmap=cmap or \"gray\", interpolation=interpolation)\n        elif img.ndim == 3:\n            # if there are 4 channels (RGBA), show first 3 (RGB). matplotlib will treat shape (H,W,3).\n            if img.shape[2] == 4:\n                ax.imshow(img[..., :3])\n            else:\n                ax.imshow(img)\n        else:\n            raise ValueError(f\"Unsupported image array shape: {img.shape}\")\n        ax.axis(\"off\")\n        if titles and idx < len(titles):\n            ax.set_title(titles[idx], fontsize=8)\n\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path, bbox_inches=\"tight\", dpi=150)\n    if show:\n        plt.show()\n    plt.close(fig)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T15:13:29.602275Z","iopub.execute_input":"2025-08-27T15:13:29.60311Z","iopub.status.idle":"2025-08-27T15:13:29.618989Z","shell.execute_reply.started":"2025-08-27T15:13:29.603071Z","shell.execute_reply":"2025-08-27T15:13:29.618226Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport pydicom\nimport numpy as np\nfrom typing import Tuple, Dict\nfrom pathlib import Path\n\n\ndef apply_dicom_windowing(img: np.ndarray, window_center: float, window_width: float) -> np.ndarray:\n    img_min = window_center - window_width // 2\n    img_max = window_center + window_width // 2\n    img = np.clip(img, img_min, img_max)\n    img = (img - img_min) / (img_max - img_min + 1e-7)\n    return img.astype(np.float32)  # Return as float32 in [0, 1]\n\n\ndef get_windowing_params(modality: str) -> Tuple[float, float]:\n    windows = {\n        'CT': (40, 80),\n        'CTA': (50, 350),\n        'MRA': (600, 1200),\n        'MRI': (40, 80),\n    }\n    return windows.get(modality, (40, 80))\n\n\ndef get_metadata_from_series_id(series_id, df):\n    sub_df = df[df['SeriesInstanceUID'] == series_id]\n    names = [\"modality\", \"sex\", \"age\"]\n    var_names = [\"Modality\", \"PatientSex\", \"PatientAge\"]\n    vars_vals = [\n        sub_df[name].iloc[0]  # Fixed: use sub_df instead of df\n        for name in var_names\n    ]\n    return dict(zip(names, vars_vals))\n\ndef process_dicom_series(\n    series_path: str,\n    df, \n    num_slices: int = 64, \n    image_size: int = 224, \n    use_windowing: bool = True) -> Tuple[np.ndarray, Dict]:\n    series_path = Path(series_path)\n    series_id = os.path.basename(series_path)\n    \n    # Find all DICOM files\n    all_filepaths = []\n    for root, _, files in os.walk(series_path):\n        for file in files:\n            if file.endswith('.dcm') or file.endswith('.DCM'):\n                all_filepaths.append(os.path.join(root, file))\n    all_filepaths.sort()\n    \n    if len(all_filepaths) == 0:\n       return None, None\n        \n    # Process DICOM files\n    slices = []\n    metadata = get_metadata_from_series_id(series_id, df)\n    \n    for i, filepath in enumerate(all_filepaths):\n        if i >= num_slices:\n            break\n        try:\n            ds = pydicom.dcmread(filepath, force=True)\n            \n            # Check if pixel data exists\n            if not hasattr(ds, 'pixel_array'):\n                print(f\"No pixel data in {filepath}\")\n                continue\n                \n            img = ds.pixel_array.astype(np.float32)\n            \n            # Handle multi-frame or color images\n            if img.ndim == 3:\n                if img.shape[-1] == 3:  # RGB image\n                    img = cv2.cvtColor(img.astype(np.uint8), cv2.COLOR_RGB2GRAY).astype(np.float32)\n                elif img.shape[0] == 3:  # Channel-first RGB\n                    img = cv2.cvtColor(img.transpose(1, 2, 0).astype(np.uint8), cv2.COLOR_RGB2GRAY).astype(np.float32)\n                else:\n                    # Multi-frame, take first frame\n                    img = img[0] if img.shape[0] < img.shape[2] else img[:, :, 0]\n            elif img.ndim > 3:\n                # Handle 4D+ data by taking first slice/frame\n                img = img[0, 0] if img.ndim == 4 else img.flatten()[:img.shape[-2]*img.shape[-1]].reshape(img.shape[-2:])\n                \n            # Apply rescale if available\n            if hasattr(ds, 'RescaleSlope') and hasattr(ds, 'RescaleIntercept'):\n                slope = float(ds.RescaleSlope) if ds.RescaleSlope != '' else 1.0\n                intercept = float(ds.RescaleIntercept) if ds.RescaleIntercept != '' else 0.0\n                img = img * slope + intercept\n            \n            # Apply windowing or normalization\n            if use_windowing:\n                window_center, window_width = get_windowing_params(metadata['modality'])\n                # Override with DICOM window values if available\n                if hasattr(ds, 'WindowCenter') and hasattr(ds, 'WindowWidth'):\n                    try:\n                        wc = ds.WindowCenter\n                        ww = ds.WindowWidth\n                        # Handle multiple window values (take first)\n                        if isinstance(wc, (list, tuple)):\n                            wc = wc[0]\n                        if isinstance(ww, (list, tuple)):\n                            ww = ww[0]\n                        window_center, window_width = float(wc), float(ww)\n                    except (ValueError, TypeError):\n                        pass  # Use default values\n                \n                img = apply_dicom_windowing(img, window_center, window_width)\n            else:\n                img_min, img_max = np.percentile(img, [1, 99])\n                if img_max > img_min:\n                    img = np.clip((img - img_min) / (img_max - img_min), 0, 1).astype(np.float32)\n                else:\n                    img = np.zeros_like(img, dtype=np.float32)\n            \n            # Resize image - ALWAYS resize to target size\n            if img.shape != (image_size, image_size):\n                img = cv2.resize(img, (image_size, image_size), interpolation=cv2.INTER_AREA)\n            \n            slices.append(img)\n            \n        except Exception as e:\n            print(f\"Error processing {filepath}: {e}\")\n            continue\n    \n    # Handle slice sampling\n    if len(slices) == 0:\n        print(\"Warning: No valid slices found!\")\n        volume = np.zeros((num_slices, image_size, image_size), dtype=np.float32)\n    else:\n        volume = np.array(slices, dtype=np.float32)\n        print(f\"Original volume shape: {volume.shape}\")\n        \n        if len(slices) > num_slices:\n            # Sample slices evenly\n            indices = np.linspace(0, len(slices) - 1, num_slices).astype(int)\n            volume = volume[indices]\n        elif len(slices) < num_slices:\n            # Pad with edge values (repeat last slice)\n            pad_size = num_slices - len(slices)\n            volume = np.pad(volume, ((0, pad_size), (0, 0), (0, 0)), mode='edge')\n    \n    return volume, metadata\n","metadata":{"trusted":true,"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import cv2\nimport numpy as np\nfrom typing import Optional, Tuple\n\n\ndef apply_non_local_means_denoising(image: np.ndarray, h: float = 10, \n                                   template_window_size: int = 7, \n                                   search_window_size: int = 21) -> np.ndarray:\n    if len(image.shape) == 2:\n        return cv2.fastNlMeansDenoising(image, None, h, template_window_size, search_window_size)\n    else:\n        return cv2.fastNlMeansDenoisingColored(image, None, h, h, template_window_size, search_window_size)\n\n\ndef apply_bilateral_filter(image: np.ndarray, d: int = 9, \n                          sigma_color: float = 75, sigma_space: float = 75) -> np.ndarray:\n    return cv2.bilateralFilter(image, d, sigma_color, sigma_space)\n\n\ndef apply_adaptive_median_filter(image: np.ndarray, max_kernel_size: int = 7) -> np.ndarray:\n    def adaptive_median_single(img, max_size):\n        result = np.copy(img)\n        rows, cols = img.shape\n        \n        for i in range(1, rows - 1):\n            for j in range(1, cols - 1):\n                kernel_size = 3\n                while kernel_size <= max_size:\n                    half_size = kernel_size // 2\n                    \n                    # Define window boundaries\n                    row_min = max(0, i - half_size)\n                    row_max = min(rows, i + half_size + 1)\n                    col_min = max(0, j - half_size)\n                    col_max = min(cols, j + half_size + 1)\n                    \n                    window = img[row_min:row_max, col_min:col_max]\n                    \n                    z_med = np.median(window)\n                    z_min = np.min(window)\n                    z_max = np.max(window)\n                    \n                    A1 = z_med - z_min\n                    A2 = z_med - z_max\n                    \n                    if A1 > 0 and A2 < 0:\n                        B1 = img[i, j] - z_min\n                        B2 = img[i, j] - z_max\n                        \n                        if B1 > 0 and B2 < 0:\n                            result[i, j] = img[i, j]\n                        else:\n                            result[i, j] = z_med\n                        break\n                    else:\n                        kernel_size += 2\n                        if kernel_size > max_size:\n                            result[i, j] = z_med\n        return result\n    \n    return adaptive_median_single(image, max_kernel_size)\n\n\ndef unsharp_mask(image: np.ndarray, kernel_size: Tuple[int, int] = (5, 5), \n                sigma: float = 1.0, amount: float = 1.0, threshold: int = 0) -> np.ndarray:\n    image_float = image.astype(np.float64)\n    \n    # Create blurred version\n    blurred = cv2.GaussianBlur(image_float, kernel_size, sigma)\n    \n    # Create sharpened image\n    sharpened = image_float + amount * (image_float - blurred)\n    \n    # Clip values to valid range\n    sharpened = np.clip(sharpened, 0, 255)\n    \n    # Apply threshold to avoid amplifying noise in low contrast areas\n    if threshold > 0:\n        low_contrast_mask = np.abs(image_float - blurred) < threshold\n        sharpened[low_contrast_mask] = image_float[low_contrast_mask]\n    \n    return sharpened.astype(np.uint8)\n\n\ndef apply_laplacian_sharpening(image: np.ndarray, alpha: float = 0.2) -> np.ndarray:\n    # Laplacian kernel for edge detection\n    kernel = np.array([[0, -1, 0], \n                       [-1, 5, -1], \n                       [0, -1, 0]], dtype=np.float32)\n    \n    # Apply Laplacian\n    laplacian = cv2.filter2D(image.astype(np.float32), -1, kernel)\n    \n    # Combine with original\n    sharpened = (1 - alpha) * image.astype(np.float32) + alpha * laplacian\n    \n    return np.clip(sharpened, 0, 255).astype(np.uint8)\n\n\ndef enhance_brain_contrast(image: np.ndarray, clip_limit: float = 2.0, \n                          tile_grid_size: Tuple[int, int] = (8, 8)) -> np.ndarray: \n    clahe = cv2.createCLAHE(clipLimit=clip_limit, tileGridSize=tile_grid_size)\n    return clahe.apply(image)\n\n\ndef enhance_brain_dicom_slice(image: np.ndarray, \n                             denoise_method: str = 'nlm',\n                             apply_sharpening: bool = True,\n                             enhance_contrast: bool = True) -> np.ndarray:\n    # Ensure image is uint8\n    if image.dtype != np.uint8:\n        image = ((image - image.min()) / (image.max() - image.min()) * 255).astype(np.uint8)\n    \n    enhanced = image.copy()\n    \n    # Step 1: Noise reduction\n    if denoise_method == 'nlm':\n        enhanced = apply_non_local_means_denoising(enhanced, h=10, template_window_size=7, search_window_size=21)\n    elif denoise_method == 'bilateral':\n        enhanced = apply_bilateral_filter(enhanced, d=9, sigma_color=75, sigma_space=75)\n    elif denoise_method == 'median':\n        enhanced = cv2.medianBlur(enhanced, 5)\n    elif denoise_method == 'adaptive_median':\n        enhanced = apply_adaptive_median_filter(enhanced, max_kernel_size=7)\n    \n    # Step 2: Contrast enhancement\n    if enhance_contrast:\n        enhanced = enhance_brain_contrast(enhanced, clip_limit=2.0, tile_grid_size=(8, 8))\n    \n    # Step 3: Sharpening\n    if apply_sharpening:\n        enhanced = unsharp_mask(enhanced, kernel_size=(5, 5), sigma=1.0, amount=0.3, threshold=5)\n    \n    return enhanced\n\n\ndef enhance_brain_dicom_series(volume: np.ndarray, \n                              denoise_method: str = 'nlm',\n                              apply_sharpening: bool = True,\n                              enhance_contrast: bool = True,\n                              progress_callback: Optional[callable] = None) -> np.ndarray:\n\n    enhanced_volume = np.zeros_like(volume)\n    num_slices = volume.shape[0]\n    \n    for i in range(num_slices):\n        enhanced_volume[i] = enhance_brain_dicom_slice(\n            volume[i], \n            denoise_method=denoise_method,\n            apply_sharpening=apply_sharpening,\n            enhance_contrast=enhance_contrast\n        )\n        \n        if progress_callback:\n            progress_callback(i + 1, num_slices)\n    \n    return enhanced_volume","metadata":{"trusted":true,"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport pandas as pd\nimport numpy as np\nfrom torch.utils.data import Dataset, DataLoader\nfrom typing import Dict, List, Optional, Tuple\nfrom pathlib import Path\nimport tqdm\nimport albumentations as A\nfrom sklearn.model_selection import train_test_split,  KFold\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nfrom data.dicom import process_dicom_series\nfrom data.image_sharpeninig import enhance_brain_dicom_series\n\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\ndef process_save_brain_ct_series(\n    to_save_dir: str,\n    root_dir: str,\n    data_df: pd.DataFrame,\n    num_slices: int = 64,\n    image_size: int = 224,\n    use_windowing: bool = True,\n    series_col: str = 'SeriesInstanceUID'\n):\n    os.makedirs(to_save_dir, exist_ok=True)\n    for _, row in tqdm.tqdm(data_df.iterrows(), desc=f\"prcoessing and saving brain ct series with: {num_slices} slices and {image_size} image size\"):\n        series_id = row[series_col]\n        series_path = os.path.join(root_dir, series_id)\n        if not os.path.exists(series_path):\n            raise FileNotFoundError(f\"Series path not found: {series_path}\")\n        \n        volume, metadata = process_dicom_series(series_path, num_slices, image_size, use_windowing)\n        if len(volume) < num_slices:\n            volume = np.concatenate([volume, np.zeros((num_slices - len(volume), image_size, image_size))])\n        data = {\n            \"volume\":volume,\n            \"metadata\":metadata\n        }\n        save_path = os.path.join(to_save_dir, f\"{series_id}.npz\")\n        np.savez(save_path, **data)\n        \n    \n\nclass RNSADataset(Dataset):\n    def __init__(\n        self,\n        data_df: pd.DataFrame,\n        root_dir: str,\n        mode: str = 'train',\n        use_windowing: bool = True,\n        enhance_images: bool = False,\n        denoise_method: str = 'nlm',\n        apply_sharpening: bool = True,\n        enhance_contrast: bool = True,\n        augmentations: Optional[A.Compose] = None,\n        series_col: str = ID_COL,\n\n    ):\n     \n        self.data_df = data_df.reset_index(drop=True)\n        self.root_dir = Path(root_dir)\n        self.mode = mode\n        self.use_windowing = use_windowing\n        self.enhance_images = enhance_images\n        self.denoise_method = denoise_method\n        self.apply_sharpening = apply_sharpening\n        self.enhance_contrast = enhance_contrast\n        self.augmentations = augmentations\n        self.series_col = series_col\n      \n    \n    def __len__(self) -> int:\n        return len(self.data_df)\n    \n    def _get_series_path(self, series_id: str) -> Path:\n        \"\"\"Get the path to a DICOM series folder\"\"\"\n        return self.root_dir / f\"{series_id}.npz\"\n    \n    def _process_metadata(self, metadata):\n        conditional_features = []\n        if metadata:\n            # Encode modality\n            modality_map = {'CT': 0, 'CTA': 1, 'MRA': 2, 'MRI': 3}\n            modality_encoded = modality_map.get(metadata.get('modality', 'CT'), 0)\n            conditional_features.append(modality_encoded)\n            \n            # Encode sex\n            sex_encoded = 1 if metadata.get('sex', 'M') == 'F' else 0\n            conditional_features.append(sex_encoded)\n            \n            # Normalize age\n            age = metadata.get('age', 50)\n            if isinstance(age, str):\n                age = int(age.replace('Y', '')) if 'Y' in age else 50\n            age_normalized = min(age / 100.0, 1.0)  # Normalize to [0, 1]\n            conditional_features.append(age_normalized)\n        else:\n            conditional_features = [0, 0, 0.5]  # Default values\n        \n        return  torch.tensor(conditional_features, dtype=torch.float32)\n\n    def set_default_metadata(self, meta):\n        self.default_metadata = meta\n    \n    def _process_volume(self, series_id: str) -> Tuple[np.ndarray, Dict]:\n        \"\"\"Process a DICOM series and return volume and metadata\"\"\"\n        series_path = self._get_series_path(series_id)\n        \n        if not series_path.exists():\n            raise FileNotFoundError(f\"Series path not found: {series_path}\")\n        data = np.load(series_path)\n        volume, metadata = data['volume'], data['metadata']\n        if self.enhance_images:\n            volume = enhance_brain_dicom_series(\n                volume,\n                denoise_method=self.denoise_method,\n                apply_sharpening=self.apply_sharpening,\n                enhance_contrast=self.enhance_contrast\n            )\n        \n        # Ensure volume is in correct format [C, D, H, W] where C=1 for grayscale\n        if volume.ndim == 3:\n            volume = volume[np.newaxis, ...]  # Add channel dimension\n        \n        return volume.astype(np.float32), metadata\n    \n    def _apply_augmentations(self, volume: np.ndarray) -> np.ndarray:\n        \"\"\"Apply augmentations to volume\"\"\"\n        if self.augmentations is None or self.mode != 'train':\n            return volume\n        \n        # Apply augmentations slice by slice\n        augmented_slices = []\n        for i in range(volume.shape[1]):  # Iterate over depth dimension\n            slice_2d = volume[0, i]  # Remove channel dim for augmentation\n            \n            # Ensure slice is in uint8 format for some augmentations\n            if slice_2d.max() <= 1.0:\n                slice_uint8 = (slice_2d * 255).astype(np.uint8)\n            else:\n                slice_uint8 = slice_2d.astype(np.uint8)\n            \n            # Apply augmentations\n            augmented = self.augmentations(image=slice_uint8)['image']\n            \n            # Convert back to float32 and normalize\n            if isinstance(augmented, torch.Tensor):\n                augmented = augmented.numpy()\n            \n            if augmented.max() > 1.0:\n                augmented = augmented.astype(np.float32) / 255.0\n            else:\n                augmented = augmented.astype(np.float32)\n            \n            augmented_slices.append(augmented)\n        \n        return np.stack(augmented_slices, axis=0)[np.newaxis, ...]  # Add channel dim back\n    \n    def _prepare_labels(self, row):\n        return [row[col_id] for col_id in LABEL_COLS]\n    \n    def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:\n        row = self.data_df.iloc[idx]\n        series_id = row[self.series_col]\n        label = self._prepare_labels(row)\n        volume, metadata = self._process_volume(series_id)\n        \n        # Check cache first\n        if not len(metadata):\n            metadata = self.default_metadata\n        \n        # Apply augmentations\n        volume = self._apply_augmentations(volume)\n        \n        # Convert to tensor\n        volume_tensor = torch.from_numpy(volume)  # Shape: [1, D, H, W]\n        label_tensor = torch.tensor(label, dtype=torch.float32)  # Float for multilabel\n        conditional_tensor = self._process_metadata(metadata)\n        \n        return {\n            'pixel_values': volume_tensor,\n            'labels': label_tensor,\n            'conditional': conditional_tensor,\n        }\n\n\ndef get_augmentations(mode: str = 'train') -> Optional[A.Compose]:\n    \"\"\"Get augmentation pipeline for different modes\"\"\"\n    if mode == 'train':\n        return A.Compose([\n            A.RandomRotate90(p=0.3),\n            A.Flip(p=0.3),\n            A.RandomBrightnessContrast(\n                brightness_limit=0.1,\n                contrast_limit=0.1,\n                p=0.3\n            ),\n            A.GaussNoise(var_limit=(0, 0.01), p=0.2),\n            A.GaussianBlur(blur_limit=(1, 3), p=0.2),\n            A.ElasticTransform(\n                alpha=50,\n                sigma=5,\n                p=0.2\n            ),\n            A.ShiftScaleRotate(\n                shift_limit=0.05,\n                scale_limit=0.05,\n                rotate_limit=5,\n                p=0.3\n            ),\n            A.Normalize(mean=[0.485], std=[0.229]),  # ImageNet normalization for grayscale\n        ])\n    else:\n        return A.Compose([\n            A.Normalize(mean=[0.485], std=[0.229]),\n        ])\n\n\ndef create_data_splits(\n    df: pd.DataFrame,\n    test_size: float = 0.1,\n    random_state: int = 42,\n) -> Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:\n    \"\"\"Create train/val/test splits\"\"\"\n    \n    train_df, test_df = train_test_split(\n        df, \n        test_size=test_size, \n        random_state=random_state\n    )\n        \n    return train_df, test_df\n\n\ndef create_cross_validation_splits(\n    df: pd.DataFrame,\n    n_splits: int = 5,\n    random_state: int = 42\n) -> List[Tuple[pd.DataFrame, pd.DataFrame]]:\n    \"\"\"Create cross-validation splits for multilabel data\"\"\"\n    \n    kf = KFold(n_splits=n_splits, shuffle=True, random_state=random_state)\n    splits = []\n    \n    for train_idx, val_idx in kf.split(df):\n        train_df = df.iloc[train_idx].reset_index(drop=True)\n        val_df = df.iloc[val_idx].reset_index(drop=True)\n        splits.append((train_df, val_df))\n    \n    print(f\"Created {n_splits}-fold cross-validation splits for multilabel data\")\n    return splits\n\n\ndef create_dataloaders(\n    train_df: pd.DataFrame,\n    val_df: pd.DataFrame,\n    root_dir: str,\n    batch_size: int = 128,\n    num_workers: int = 2,\n    enhance_images: bool = False,\n    aug=False,\n    **dataset_kwargs\n) -> Dict[str, DataLoader]:\n    \"\"\"Create train/val/test dataloaders\"\"\"\n    \n    # Get augmentations\n    train_augs = None\n    val_augs = None\n    if aug:\n       train_augs = get_augmentations('train')\n       val_augs = get_augmentations('val')\n    \n    # Create datasets\n    train_dataset = RNSADataset(\n        train_df,\n        root_dir,\n        mode='train',\n        enhance_images=enhance_images,\n        augmentations=train_augs,\n        **dataset_kwargs\n    )\n    \n    val_dataset = RNSADataset(\n        val_df,\n        root_dir,\n        mode='val',\n        enhance_images=enhance_images,\n        augmentations=val_augs,\n        **dataset_kwargs\n    )\n    \n    # Create dataloaders\n    dataloaders = {\n        'train': DataLoader(\n            train_dataset,\n            batch_size=batch_size,\n            shuffle=True,\n            num_workers=num_workers,\n            pin_memory=True,\n            drop_last=True\n        ),\n        'val': DataLoader(\n            val_dataset,\n            batch_size=batch_size,\n            shuffle=False,\n            num_workers=num_workers,\n            pin_memory=True,\n            drop_last=False\n        )\n    }\n    return dataloaders\n","metadata":{"trusted":true,"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch import nn\nimport torch\nfrom typing import List\nimport torch.nn.functional as F\n\n\n\ndef normalize(in_channels: int):\n    return torch.nn.GroupNorm(num_groups=32,\n                              num_channels=in_channels,\n                              eps=1e-6, affine=True)\n\n\ndef nonlinearity(name: str):\n    return getattr(nn, name)()\n\nclass ResidualBlock(nn.Module):\n    def __init__(\n        self,\n        in_channels: int,\n        out_channels: int,\n        activation: str = 'SiLU',\n        dropout: float = 0.2,\n        attn=False,\n        conditional_dim = None\n    ):\n        super().__init__()\n        \n        self.block1 = nn.Sequential(\n            normalize(in_channels),\n            nonlinearity(activation),\n            nn.Conv2d(in_channels, out_channels, 3, padding=1)\n        )\n        \n        self.block2 = nn.Sequential(\n            normalize(out_channels),\n            nonlinearity(activation),\n            nn.Dropout(dropout),\n            nn.Conv2d(out_channels, out_channels, 3, padding=1)\n        )\n        \n        self.shortcut = nn.Identity()\n        if in_channels != out_channels:\n            self.shortcut = nn.Conv2d(in_channels, out_channels, 1)\n\n        if attn:\n          self.attn = AttentionBlock(out_channels)\n        else:\n          self.attn = nn.Identity()\n\n        if conditional_dim is not None:\n           self.cond_proj = nn.Sequential(\n               nonlinearity(activation),\n               nn.Linear(conditional_dim, out_channels)\n           )\n\n    def forward(self, x: torch.Tensor, cond_embd = None):\n        h = self.block1(x)\n        if cond_embd is not None:\n            cond_embd = self.cond_proj(cond_embd)\n            h = h + cond_embd[:, :, None, None] \n        h = self.block2(h)\n        return self.attn(self.shortcut(x) + h)\n\n\nclass AttentionBlock(nn.Module):\n    def __init__(self, in_channels: int):\n        super().__init__()\n        self.in_channels = in_channels\n        self.norm = normalize(in_channels)\n        self.q = nn.Conv2d(in_channels, in_channels, 1)\n        self.k = nn.Conv2d(in_channels, in_channels, 1)\n        self.v = nn.Conv2d(in_channels, in_channels, 1)\n        self.proj_out = nn.Conv2d(in_channels, in_channels, 1)\n\n    def forward(self, x: torch.Tensor):\n        h = self.norm(x)\n        q = self.q(h)\n        k = self.k(h)\n        v = self.v(h)\n\n        # Reshape for attention computation\n        b, c, h, w = q.shape\n        q = q.view(b, c, h * w).transpose(1, 2)  # b, hw, c\n        k = k.view(b, c, h * w).transpose(1, 2)  # b, hw, c\n        v = v.view(b, c, h * w).transpose(1, 2)  # b, hw, c\n        \n        # Use scaled dot product attention\n        attn_output = F.scaled_dot_product_attention(q, k, v)\n        \n        # Reshape back\n        attn_output = attn_output.transpose(1, 2).view(b, c, h, w)\n        h = self.proj_out(attn_output)\n        \n        return x + h\n\n\nclass DownsampleBlock(nn.Module):\n    def __init__(self, channels: int, use_conv: bool = True):\n        super().__init__()\n        if use_conv:\n            self.downsample = nn.Conv2d(channels, channels, 3, stride=2, padding=1)\n        else:\n            self.downsample = nn.AvgPool2d(2, 2)\n\n    def forward(self, x: torch.Tensor):\n        return self.downsample(x)\n\n\nclass UpsampleBlock(nn.Module):\n    def __init__(self, channels: int, use_conv: bool = True, mode: str = 'nearest'):\n        super().__init__()\n        self.use_conv = use_conv\n        self.mode = mode\n        if use_conv:\n            self.upsample = nn.Conv2d(channels, channels, 3, padding=1)\n\n    def forward(self, x: torch.Tensor):\n        x = F.interpolate(x, scale_factor=2.0, mode=self.mode)\n        if self.use_conv:\n            x = self.upsample(x)\n        return x\n\n\n","metadata":{"trusted":true,"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch import nn\nimport torch\nimport torch.nn.functional as F\nfrom typing import List, Optional, Callable, Union, Any\nfrom modules import (\n    ResidualBlock,\n    UpsampleBlock,\n    DownsampleBlock,\n    AttentionBlock,\n    nonlinearity,\n    normalize\n)\nfrom dataclasses import dataclass\nfrom transformers import AutoModel\nimport math\n\n\n@dataclass\nclass UNetConfig:\n    in_channels: int = 3\n    out_channels: int = 3\n    base_channels: int = 32\n    channel_multipliers: List[int] = None\n    attention_resolutions: List[int] = None\n    block_depth: int = 2\n    activation: str = 'SiLU'\n    use_attention: bool = True\n    use_conv_in_downsample: bool = True\n    use_conv_in_upsample: bool = True\n    dropout: float = 0.3\n    input_resolution: int = 224\n    cond_input_dim: Optional[int] = 2\n    cond_embd_dim: Optional[int] = 128\n    \n    def __post_init__(self):\n        if self.channel_multipliers is None:\n            self.channel_multipliers = [2, 2, 4, 4, 8]\n        if self.attention_resolutions is None:\n            self.attention_resolutions = [28, 14, 7] \n\n\n@dataclass\nclass MainConfig:\n    num_slices: int = 64\n    base_model_id: str = \"facebook/convnextv2-base-22k-224\"\n    trust_remote_code: bool = True\n    num_labels: int = 2\n    head_type: str = \"mlp\"   \n    head_activation: Union[str, Callable[[torch.Tensor], torch.Tensor]] = \"gelu\"\n    head_dropout: float = 0.2\n    head_hidden_sizes: List[int] = None\n    pooling: str = \"cls\"\n    freeze_base: bool = True\n    feature_dim: Optional[int] = None\n    \n    def __post_init__(self):\n        if self.head_hidden_sizes is None:\n            self.head_hidden_sizes = [512, 256]\n    \n\ndef _get_activation(act: Union[str, Callable]):\n    if callable(act):\n        return act\n    act = act.lower()\n    if act in (\"relu\",):\n        return nn.ReLU()\n    if act in (\"gelu\",):\n        return nn.GELU()\n    if act in (\"tanh\",):\n        return nn.Tanh()\n    if act in (\"swish\", \"silu\"):\n        return nn.SiLU()\n    return nn.GELU()\n\nclass UNetFeatureExtractor(nn.Module):\n    \"\"\"UNet-based feature extractor that outputs feature maps instead of reconstructed images\"\"\"\n    \n    def __init__(self, config: UNetConfig):\n        super().__init__()\n        self.config = config\n        \n        # Input projection\n        self.conv_in = nn.Conv2d(config.in_channels, config.base_channels, 3, padding=1)\n        \n        # Conditional embedding\n        if config.cond_input_dim is not None:\n            self.cond_proj = nn.Sequential(\n                nn.Linear(config.cond_input_dim, config.cond_embd_dim),\n                nonlinearity(config.activation),\n                nn.Linear(config.cond_embd_dim, config.cond_embd_dim)\n            )\n        else:\n            self.cond_proj = None\n        \n        # Build channel dimensions\n        ch_mult = [1] + config.channel_multipliers\n        \n        downs = []\n        down_channels = []\n        current_res = config.input_resolution\n        \n        for i in range(len(config.channel_multipliers)):\n            in_ch = config.base_channels * ch_mult[i]\n            out_ch = config.base_channels * ch_mult[i + 1]\n            \n            res_blocks = []\n            for j in range(config.block_depth):\n                block_in_ch = in_ch if j == 0 else out_ch\n                res_blocks.append(\n                    ResidualBlock(block_in_ch, out_ch, config.activation, config.dropout, conditional_dim=config.cond_embd_dim)\n                )\n            \n            # Attention block (optional)\n            use_attn = config.use_attention and current_res in config.attention_resolutions\n            attn_block = AttentionBlock(out_ch) if use_attn else nn.Identity()\n            \n            # Downsample\n            downsample = DownsampleBlock(out_ch, config.use_conv_in_downsample)\n            \n            down_block = nn.Module()\n            down_block.res_blocks = nn.ModuleList(res_blocks)\n            down_block.attn = attn_block\n            down_block.downsample = downsample\n            \n            downs.append(down_block)\n            down_channels.append(out_ch)\n            current_res //= 2\n            \n        self.downs = nn.ModuleList(downs)\n        \n        # Middle blocks\n        mid_ch = config.base_channels * config.channel_multipliers[-1]\n        use_mid_attn = config.use_attention and current_res in config.attention_resolutions\n        self.mid_block1 = ResidualBlock(mid_ch, mid_ch, config.activation, config.dropout, attn=use_mid_attn, conditional_dim=config.cond_embd_dim)\n        self.mid_block2 = ResidualBlock(mid_ch, mid_ch, config.activation, config.dropout, attn=use_mid_attn, conditional_dim=config.cond_embd_dim)\n        \n        # Output projection to produce 2D images for base model\n        self.output_proj = nn.Sequential(\n            normalize(mid_ch),\n            nonlinearity(config.activation),\n            nn.Conv2d(mid_ch, config.out_channels, 3, padding=1)\n        )\n\n    def forward(\n        self,\n        x: torch.Tensor,\n        cond: Optional[torch.Tensor] = None\n    ):\n        # Handle 3D input [B, C, D, H, W] by processing slice by slice\n        if x.ndim == 5:\n            batch_size, channels, depth, height, width = x.shape\n            # Reshape to [B*D, C, H, W]\n            x = x.view(batch_size * depth, channels, height, width)\n            process_3d = True\n        else:\n            batch_size, depth = x.shape[0], 1\n            process_3d = False\n        \n        # Process conditional embedding\n        if cond is not None and self.cond_proj is not None:\n            if process_3d:\n                # Repeat conditional for each slice\n                cond = cond.unsqueeze(1).repeat(1, depth, 1).view(batch_size * depth, -1)\n            cond_embd = self.cond_proj(cond)\n        else:\n            cond_embd = None\n        \n        # Forward pass through UNet encoder\n        h = self.conv_in(x)\n        \n        for down_block in self.downs:\n            for res_block in down_block.res_blocks:\n                h = res_block(h, cond_embd)\n            h = down_block.attn(h)\n            h = down_block.downsample(h)\n        \n        # Middle blocks\n        h = self.mid_block1(h, cond_embd)\n        h = self.mid_block2(h, cond_embd)\n        \n        # Generate output images\n        output = self.output_proj(h)  # [B*D, C, H, W] or [B, C, H, W]\n        \n        if process_3d:\n            # Reshape back to [B, D, C, H, W] and take middle slice\n            output = output.view(batch_size, depth, *output.shape[1:])\n            # Take middle slice for base model input\n            middle_idx = depth // 2\n            output = output[:, middle_idx]  # [B, C, H, W]\n        \n        return output\n\n\nclass UNet(nn.Module):\n    \"\"\"Original UNet implementation for reconstruction tasks\"\"\"\n    \n    def __init__(self, config: UNetConfig):\n        super().__init__()\n        self.config = config\n        # Input projection\n        self.conv_in = nn.Conv2d(config.in_channels, config.base_channels, 3, padding=1)\n        if config.cond_input_dim is not None:\n            self.cond_proj = nn.Sequential(\n                nn.Linear(config.cond_input_dim, config.cond_embd_dim),\n                nonlinearity(config.activation),\n                nn.Linear(config.cond_embd_dim, config.cond_embd_dim)\n            )\n        \n        # Build channel dimensions\n        ch_mult = [1] + config.channel_multipliers\n        \n        downs = []\n        down_channels = []\n        current_res = config.input_resolution\n        \n        for i in range(len(config.channel_multipliers)):\n            in_ch = config.base_channels * ch_mult[i]\n            out_ch = config.base_channels * ch_mult[i + 1]\n            \n            res_blocks = []\n            for j in range(config.block_depth):\n                block_in_ch = in_ch if j == 0 else out_ch\n                res_blocks.append(\n                    ResidualBlock(block_in_ch, out_ch, config.activation, config.dropout, conditional_dim=config.cond_embd_dim)\n                )\n            \n            # Attention block (optional)\n            use_attn = config.use_attention and current_res in config.attention_resolutions\n            attn_block = AttentionBlock(out_ch) if use_attn else nn.Identity()\n            \n            # Downsample\n            downsample = DownsampleBlock(out_ch, config.use_conv_in_downsample)\n            \n            down_block = nn.Module()\n            down_block.res_blocks = nn.ModuleList(res_blocks)\n            down_block.attn = attn_block\n            down_block.downsample = downsample\n            \n            downs.append(down_block)\n            down_channels.append(out_ch)\n            current_res //= 2\n            \n        self.downs = nn.ModuleList(downs)\n        \n        # Middle blocks\n        mid_ch = config.base_channels * config.channel_multipliers[-1]\n        use_mid_attn = config.use_attention and current_res in config.attention_resolutions\n        self.mid_block1 = ResidualBlock(mid_ch, mid_ch, config.activation, config.dropout, attn=use_mid_attn, conditional_dim=config.cond_embd_dim)\n        self.mid_block2 = ResidualBlock(mid_ch, mid_ch, config.activation, config.dropout, attn=use_mid_attn, conditional_dim=config.cond_embd_dim)\n        \n        # Upsampling path\n        ups = []\n        current_res = config.input_resolution // (2 ** len(config.channel_multipliers))\n        \n        for i in reversed(range(len(config.channel_multipliers))):\n            current_res *= 2\n            out_ch = config.base_channels * ch_mult[i + 1]\n            skip_ch = down_channels[i]\n            \n            if i == len(config.channel_multipliers) - 1:\n                upsample_in_ch = mid_ch\n            else:\n                upsample_in_ch = config.base_channels * ch_mult[i + 2]\n            \n            upsample = UpsampleBlock(upsample_in_ch, config.use_conv_in_upsample)\n            \n            # Residual blocks\n            res_blocks = []\n            for j in range(config.block_depth):\n                if j == 0:\n                    block_in_ch = upsample_in_ch + skip_ch\n                    block_out_ch = out_ch\n                else:\n                    block_in_ch = out_ch\n                    block_out_ch = out_ch\n                res_blocks.append(\n                    ResidualBlock(block_in_ch, block_out_ch, config.activation, config.dropout, conditional_dim=config.cond_embd_dim))\n            \n            # Attention block (optional)\n            use_attn = config.use_attention and current_res in config.attention_resolutions\n            attn_block = AttentionBlock(out_ch) if use_attn else nn.Identity()\n            \n            up_block = nn.Module()\n            up_block.upsample = upsample\n            up_block.res_blocks = nn.ModuleList(res_blocks)\n            up_block.attn = attn_block\n            \n            ups.append(up_block)\n            \n        self.ups = nn.ModuleList(ups)\n        \n        # Output projection\n        final_ch = config.base_channels * ch_mult[1]\n        self.conv_out = nn.Sequential(\n            normalize(final_ch),\n            nonlinearity(config.activation),\n            nn.Conv2d(final_ch, config.out_channels, 3, padding=1)\n        )\n\n    def forward(\n        self,\n        x: torch.Tensor,\n        cond: Optional[torch.Tensor] = None\n    ):\n        h = self.conv_in(x)\n        if cond is not None:\n            cond = self.cond_proj(cond)\n        # Store skip connections\n        skip_connections = []\n        \n        for down_block in self.downs:\n            for res_block in down_block.res_blocks:\n                h = res_block(h, cond)\n            h = down_block.attn(h)\n            skip_connections.append(h)\n            h = down_block.downsample(h)\n        \n        h = self.mid_block1(h, cond)\n        h = self.mid_block2(h, cond)\n        \n        for up_block in self.ups:\n            h = up_block.upsample(h)\n            skip = skip_connections.pop()\n            h = torch.cat([h, skip], dim=1)\n            \n            for res_block in up_block.res_blocks:\n                h = res_block(h, cond)\n        \n            h = up_block.attn(h)\n        \n        return self.conv_out(h)\n    \n\ndef get_unet(**config_params):\n    config = UNetConfig(**config_params)\n    return UNet(config)\n\n\ndef get_unet_feature_extractor(**config_params):\n    \"\"\"Create UNet feature extractor for preprocessing volumes\"\"\"\n    config = UNetConfig(**config_params)\n    return UNetFeatureExtractor(config)\n\n\n\nclass MainModel(nn.Module):\n    def __init__(\n        self, \n        config: MainConfig, \n        feature_ext=None,\n        custom_head: Optional[nn.Module] = None\n    ):\n        super().__init__()\n        self.config = config\n        self.feature_ext = feature_ext\n        \n        # Initialize base model\n        self.base_model = AutoModel.from_pretrained(\n            config.base_model_id, \n            trust_remote_code=config.trust_remote_code\n        )\n\n        self.set_base_model_training(config.freeze_base)\n\n        # Infer base model feature dimension\n        self.base_feature_dim = self._infer_feature_dim(config)\n        if self.base_feature_dim is None:\n            raise ValueError(\n                \"Could not infer feature dim from base model config. \"\n                \"Set config.feature_dim to the correct embedding size (e.g. 768).\"\n            )\n\n        # Build classifier head\n        self.num_labels = config.num_labels\n        if config.head_type == \"custom\":\n            if custom_head is None:\n                raise ValueError(\"custom_head must be provided when config.head_type == 'custom'\")\n            self.classifier = custom_head\n        elif config.head_type == \"linear\":\n            self.classifier = nn.Sequential(\n                nn.Dropout(config.head_dropout),\n                nn.Linear(self.base_feature_dim, self.num_labels)\n            )\n            self._init_linear(self.classifier[-1])\n        elif config.head_type == \"mlp\":\n            layers = []\n            in_dim = self.base_feature_dim\n            act_module = _get_activation(config.head_activation)\n            for h in config.head_hidden_sizes:\n                layers.append(nn.Linear(in_dim, h))\n                layers.append(act_module)\n                if config.head_dropout and config.head_dropout > 0:\n                    layers.append(nn.Dropout(config.head_dropout))\n                in_dim = h\n            layers.append(nn.Linear(in_dim, self.num_labels))\n            self.classifier = nn.Sequential(*layers)\n            # Init last linear\n            self._init_linear(self.classifier[-1])\n        else:\n            raise ValueError(f\"Unknown head_type: {config.head_type}\")\n        \n    def set_base_model_training(self, training):\n        for p in self.base_model.parameters():\n            p.requires_grad = training\n\n    def _infer_feature_dim(self, config: MainConfig) -> Optional[int]:\n        if config.feature_dim is not None:\n            return config.feature_dim\n\n        mc = getattr(self.base_model, \"config\", None)\n        if mc is None:\n            return None\n        for attr in (\"hidden_size\", \"embed_dim\", \"d_model\", \"feature_size\", \"num_features\"):\n            if hasattr(mc, attr):\n                val = getattr(mc, attr)\n                if isinstance(val, int):\n                    return val\n        if hasattr(mc, \"hidden_sizes\") and isinstance(mc.hidden_sizes, (list, tuple)) and len(mc.hidden_sizes) > 0:\n            last = mc.hidden_sizes[-1]\n            if isinstance(last, int):\n                return last\n        return None\n\n    def _init_linear(self, layer: nn.Module):\n        if isinstance(layer, nn.Linear):\n            nn.init.xavier_uniform_(layer.weight)\n            if layer.bias is not None:\n                nn.init.zeros_(layer.bias)\n\n    def _get_pooled_features(self, outputs: Any, pooling: str = \"cls\"):\n        if hasattr(outputs, \"pooler_output\"):\n            return outputs.pooler_output\n        if hasattr(outputs, \"pooled_output\"):\n            return outputs.pooled_output\n        last = None\n        if hasattr(outputs, \"last_hidden_state\"):\n            last = outputs.last_hidden_state\n        elif isinstance(outputs, (tuple, list)) and len(outputs) > 0:\n            candidate = outputs[0]\n            if isinstance(candidate, torch.Tensor) and candidate.ndim == 3:\n                last = candidate\n\n        if last is not None:\n            if pooling == \"mean\":\n                return last.mean(dim=1)\n            else:\n                return last[:, 0, :]\n\n        raise RuntimeError(\"Unable to extract features from base model outputs.\")\n\n    def forward(\n        self,\n        pixel_values: Optional[torch.Tensor] = None,\n        cond: Optional[torch.Tensor] = None,\n        **kwargs\n    ):\n        # Process input for base model\n        pixel_values = self.feature_ext(pixel_values, cond)\n      \n        model_inputs = {\"pixel_values\": pixel_values}\n        model_inputs.update(kwargs)\n        outputs = self.base_model(**model_inputs)\n        features = self._get_pooled_features(outputs, pooling=self.config.pooling)\n       \n        # Classification\n        logits = self.classifier(features)\n        return logits\n","metadata":{"trusted":true,"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom typing import Dict, List, Optional, Tuple, Union\nfrom sklearn.metrics import (\n    accuracy_score, \n    precision_score, \n    recall_score, \n    f1_score,\n    roc_auc_score, \n    roc_curve,\n    precision_recall_curve,\n    confusion_matrix,\n    classification_report\n)\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom pathlib import Path\nimport json\nimport pickle\nfrom dataclasses import dataclass, asdict\nfrom collections import defaultdict\nimport warnings\nwarnings.filterwarnings('ignore')\n\n\n@dataclass\nclass MetricResults:\n    \"\"\"Container for metric results\"\"\"\n    accuracy: float\n    precision: float\n    recall: float\n    f1: float\n    auc_roc: float\n    auc_pr: float\n    loss: float\n    \n    def to_dict(self) -> Dict[str, float]:\n        return asdict(self)\n    \n    def __str__(self) -> str:\n        return (f\"Accuracy: {self.accuracy:.4f}, Precision: {self.precision:.4f}, \"\n                f\"Recall: {self.recall:.4f}, F1: {self.f1:.4f}, \"\n                f\"AUC-ROC: {self.auc_roc:.4f}, AUC-PR: {self.auc_pr:.4f}, \"\n                f\"Loss: {self.loss:.4f}\")\n\n\nclass MetricsTracker:\n    \"\"\"Comprehensive metrics tracking for binary, multiclass, and multilabel classification\"\"\"\n    \n    def __init__(\n        self,\n        num_classes: int = 2,\n        average: str = 'macro',\n        task_type: str = 'multilabel',  # 'binary', 'multiclass', or 'multilabel'\n        threshold: float = 0.5,  # Threshold for multilabel predictions\n        save_predictions: bool = True,\n        save_plots: bool = True\n    ):\n        self.num_classes = num_classes\n        self.average = average if task_type != 'binary' else 'binary'\n        self.task_type = task_type\n        self.threshold = threshold\n        self.save_predictions = save_predictions\n        self.save_plots = save_plots\n        \n        # Storage for predictions and targets\n        self.reset()\n        \n        # History for tracking over epochs\n        self.train_history = defaultdict(list)\n        self.val_history = defaultdict(list)\n        self.test_history = defaultdict(list)\n    \n    def reset(self):\n        \"\"\"Reset accumulated predictions and targets\"\"\"\n        self.predictions = []\n        self.probabilities = []\n        self.targets = []\n        self.losses = []\n    \n    def update(\n        self, \n        logits: torch.Tensor, \n        targets: torch.Tensor, \n        loss: Optional[torch.Tensor] = None\n    ):\n        \"\"\"Update with batch predictions and targets\"\"\"\n        # Convert to numpy\n        if isinstance(logits, torch.Tensor):\n            logits = logits.detach().cpu()\n        if isinstance(targets, torch.Tensor):\n            targets = targets.detach().cpu()\n        if loss is not None and isinstance(loss, torch.Tensor):\n            loss = loss.detach().cpu()\n        \n        # Get probabilities and predictions based on task type\n        if self.task_type == 'multilabel':\n            # Multilabel classification\n            probs = torch.sigmoid(logits).numpy()  # Shape: [batch_size, num_classes]\n            preds = (probs >= self.threshold).astype(int)\n            \n            # Store as lists for each sample\n            self.predictions.extend(preds.tolist())\n            self.probabilities.extend(probs.tolist())\n            self.targets.extend(targets.numpy().tolist())\n            \n        elif self.task_type == 'binary':\n            if logits.shape[-1] == 1:\n                # Single output for binary classification\n                probs = torch.sigmoid(logits).squeeze(-1).numpy()\n                preds = (probs >= 0.5).astype(int)\n            else:\n                # Two outputs for binary classification\n                probs = torch.softmax(logits, dim=-1)[:, 1].numpy()  # Positive class probability\n                preds = torch.argmax(logits, dim=-1).numpy()\n            \n            self.predictions.extend(preds)\n            self.probabilities.extend(probs)\n            self.targets.extend(targets.numpy())\n            \n        else:\n            # Multiclass\n            probs = torch.softmax(logits, dim=-1).numpy()\n            preds = torch.argmax(logits, dim=-1).numpy()\n            \n            self.predictions.extend(preds)\n            self.probabilities.extend(probs)\n            self.targets.extend(targets.numpy())\n        \n        if loss is not None:\n            self.losses.extend([loss.item()] * len(targets))\n    \n    def compute_metrics(self) -> MetricResults:\n        \"\"\"Compute all metrics from accumulated predictions\"\"\"\n        if len(self.predictions) == 0:\n            raise ValueError(\"No predictions to compute metrics from\")\n        \n        predictions = np.array(self.predictions)\n        targets = np.array(self.targets)\n        \n        if self.task_type == 'multilabel':\n            # Multilabel classification metrics\n            # Subset accuracy (exact match)\n            accuracy = accuracy_score(targets, predictions)\n            \n            # Micro/macro averaged metrics\n            precision = precision_score(targets, predictions, average=self.average, zero_division=0)\n            recall = recall_score(targets, predictions, average=self.average, zero_division=0)\n            f1 = f1_score(targets, predictions, average=self.average, zero_division=0)\n            \n            # AUC metrics for multilabel\n            probabilities = np.array(self.probabilities)\n            try:\n                # Check if we have valid probabilities for AUC calculation\n                if targets.sum() > 0 and targets.sum() < targets.size:  # Not all zeros or all ones\n                    auc_roc = roc_auc_score(targets, probabilities, average=self.average)\n                else:\n                    auc_roc = 0.0\n                \n                # Average precision (area under PR curve) for multilabel\n                from sklearn.metrics import average_precision_score\n                auc_pr = average_precision_score(targets, probabilities, average=self.average)\n            except (ValueError, Exception):\n                auc_roc = 0.0\n                auc_pr = 0.0\n                \n        elif self.task_type == 'binary':\n            accuracy = accuracy_score(targets, predictions)\n            precision = precision_score(targets, predictions, average='binary', zero_division=0)\n            recall = recall_score(targets, predictions, average='binary', zero_division=0)\n            f1 = f1_score(targets, predictions, average='binary', zero_division=0)\n            \n            # AUC metrics\n            probabilities = np.array(self.probabilities)\n            if len(np.unique(targets)) > 1:  # Ensure both classes are present\n                auc_roc = roc_auc_score(targets, probabilities)\n                precision_curve, recall_curve, _ = precision_recall_curve(targets, probabilities)\n                auc_pr = np.trapz(precision_curve, recall_curve)\n            else:\n                auc_roc = 0.0\n                auc_pr = 0.0\n                \n        else:\n            # Multiclass\n            accuracy = accuracy_score(targets, predictions)\n            precision = precision_score(targets, predictions, average=self.average, zero_division=0)\n            recall = recall_score(targets, predictions, average=self.average, zero_division=0)\n            f1 = f1_score(targets, predictions, average=self.average, zero_division=0)\n            \n            # Multiclass AUC\n            if len(np.unique(targets)) > 1:\n                probabilities = np.array(self.probabilities)\n                if probabilities.ndim == 1:\n                    # If only positive class probabilities stored, can't compute multiclass AUC\n                    auc_roc = 0.0\n                else:\n                    auc_roc = roc_auc_score(targets, probabilities, multi_class='ovr', average=self.average)\n                auc_pr = 0.0  # PR AUC not typically used for multiclass\n            else:\n                auc_roc = 0.0\n                auc_pr = 0.0\n        \n        # Average loss\n        loss = np.mean(self.losses) if self.losses else 0.0\n        \n        return MetricResults(\n            accuracy=accuracy,\n            precision=precision,\n            recall=recall,\n            f1=f1,\n            auc_roc=auc_roc,\n            auc_pr=auc_pr,\n            loss=loss\n        )\n    \n    def save_epoch_metrics(self, metrics: MetricResults, phase: str = 'train'):\n        \"\"\"Save metrics for an epoch\"\"\"\n        history = getattr(self, f\"{phase}_history\")\n        for key, value in metrics.to_dict().items():\n            history[key].append(value)\n    \n    def get_best_metrics(self, phase: str = 'val', metric: str = 'auc_roc') -> Tuple[MetricResults, int]:\n        \"\"\"Get best metrics and epoch for a given phase and metric\"\"\"\n        history = getattr(self, f\"{phase}_history\")\n        if metric not in history or len(history[metric]) == 0:\n            raise ValueError(f\"No {metric} history found for {phase}\")\n        \n        if metric in ['loss']:\n            # Lower is better\n            best_idx = np.argmin(history[metric])\n        else:\n            # Higher is better\n            best_idx = np.argmax(history[metric])\n        \n        best_metrics = MetricResults(\n            accuracy=history['accuracy'][best_idx],\n            precision=history['precision'][best_idx],\n            recall=history['recall'][best_idx],\n            f1=history['f1'][best_idx],\n            auc_roc=history['auc_roc'][best_idx],\n            auc_pr=history['auc_pr'][best_idx],\n            loss=history['loss'][best_idx]\n        )\n        \n        return best_metrics, best_idx\n    \n    def plot_training_curves(self, save_path: Optional[str] = None, show: bool = True):\n        \"\"\"Plot training and validation curves\"\"\"\n        metrics_to_plot = ['loss', 'accuracy', 'auc_roc', 'f1']\n        \n        fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n        axes = axes.flatten()\n        \n        for i, metric in enumerate(metrics_to_plot):\n            ax = axes[i]\n            \n            if metric in self.train_history and len(self.train_history[metric]) > 0:\n                epochs = range(1, len(self.train_history[metric]) + 1)\n                ax.plot(epochs, self.train_history[metric], 'b-', label=f'Train {metric}', linewidth=2)\n            \n            if metric in self.val_history and len(self.val_history[metric]) > 0:\n                epochs = range(1, len(self.val_history[metric]) + 1)\n                ax.plot(epochs, self.val_history[metric], 'r-', label=f'Val {metric}', linewidth=2)\n            \n            ax.set_title(f'{metric.replace(\"_\", \" \").title()} Over Epochs')\n            ax.set_xlabel('Epoch')\n            ax.set_ylabel(metric.replace(\"_\", \" \").title())\n            ax.legend()\n            ax.grid(True, alpha=0.3)\n        \n        plt.tight_layout()\n        \n        if save_path:\n            plt.savefig(save_path, dpi=300, bbox_inches='tight')\n            print(f\"Training curves saved to {save_path}\")\n        \n        if show:\n            plt.show()\n        else:\n            plt.close()\n    \n    def plot_confusion_matrix(self, save_path: Optional[str] = None, show: bool = True):\n        \"\"\"Plot confusion matrix\"\"\"\n        if len(self.predictions) == 0:\n            print(\"No predictions available for confusion matrix\")\n            return\n        \n        predictions = np.array(self.predictions)\n        targets = np.array(self.targets)\n        \n        cm = confusion_matrix(targets, predictions)\n        \n        plt.figure(figsize=(8, 6))\n        sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n                    xticklabels=range(self.num_classes),\n                    yticklabels=range(self.num_classes))\n        plt.title('Confusion Matrix')\n        plt.xlabel('Predicted')\n        plt.ylabel('Actual')\n        \n        if save_path:\n            plt.savefig(save_path, dpi=300, bbox_inches='tight')\n            print(f\"Confusion matrix saved to {save_path}\")\n        \n        if show:\n            plt.show()\n        else:\n            plt.close()\n    \n    def plot_roc_curve(self, save_path: Optional[str] = None, show: bool = True):\n        \"\"\"Plot ROC curve for binary classification\"\"\"\n        if self.task_type != 'binary' or len(self.predictions) == 0:\n            print(\"ROC curve only available for binary classification with predictions\")\n            return\n        \n        targets = np.array(self.targets)\n        probabilities = np.array(self.probabilities)\n        \n        if len(np.unique(targets)) < 2:\n            print(\"ROC curve requires both classes to be present\")\n            return\n        \n        fpr, tpr, _ = roc_curve(targets, probabilities)\n        auc = roc_auc_score(targets, probabilities)\n        \n        plt.figure(figsize=(8, 6))\n        plt.plot(fpr, tpr, 'b-', linewidth=2, label=f'ROC Curve (AUC = {auc:.3f})')\n        plt.plot([0, 1], [0, 1], 'r--', linewidth=1, label='Random Classifier')\n        plt.xlabel('False Positive Rate')\n        plt.ylabel('True Positive Rate')\n        plt.title('Receiver Operating Characteristic (ROC) Curve')\n        plt.legend()\n        plt.grid(True, alpha=0.3)\n        \n        if save_path:\n            plt.savefig(save_path, dpi=300, bbox_inches='tight')\n            print(f\"ROC curve saved to {save_path}\")\n        \n        if show:\n            plt.show()\n        else:\n            plt.close()\n    \n    def plot_precision_recall_curve(self, save_path: Optional[str] = None, show: bool = True):\n        \"\"\"Plot precision-recall curve for binary classification\"\"\"\n        if self.task_type != 'binary' or len(self.predictions) == 0:\n            print(\"PR curve only available for binary classification with predictions\")\n            return\n        \n        targets = np.array(self.targets)\n        probabilities = np.array(self.probabilities)\n        \n        if len(np.unique(targets)) < 2:\n            print(\"PR curve requires both classes to be present\")\n            return\n        \n        precision, recall, _ = precision_recall_curve(targets, probabilities)\n        auc_pr = np.trapz(precision, recall)\n        \n        plt.figure(figsize=(8, 6))\n        plt.plot(recall, precision, 'b-', linewidth=2, label=f'PR Curve (AUC = {auc_pr:.3f})')\n        plt.xlabel('Recall')\n        plt.ylabel('Precision')\n        plt.title('Precision-Recall Curve')\n        plt.legend()\n        plt.grid(True, alpha=0.3)\n        \n        if save_path:\n            plt.savefig(save_path, dpi=300, bbox_inches='tight')\n            print(f\"PR curve saved to {save_path}\")\n        \n        if show:\n            plt.show()\n        else:\n            plt.close()\n    \n    def save_predictions(self, save_path: str):\n        \"\"\"Save predictions and targets\"\"\"\n        data = {\n            'predictions': self.predictions,\n            'probabilities': self.probabilities,\n            'targets': self.targets,\n            'losses': self.losses\n        }\n        \n        with open(save_path, 'wb') as f:\n            pickle.dump(data, f)\n        print(f\"Predictions saved to {save_path}\")\n    \n    def load_predictions(self, load_path: str):\n        \"\"\"Load predictions and targets\"\"\"\n        with open(load_path, 'rb') as f:\n            data = pickle.load(f)\n        \n        self.predictions = data['predictions']\n        self.probabilities = data['probabilities']\n        self.targets = data['targets']\n        self.losses = data.get('losses', [])\n        print(f\"Predictions loaded from {load_path}\")\n    \n    def save_history(self, save_path: str):\n        \"\"\"Save training history\"\"\"\n        history = {\n            'train': dict(self.train_history),\n            'val': dict(self.val_history),\n            'test': dict(self.test_history)\n        }\n        \n        with open(save_path, 'w') as f:\n            json.dump(history, f, indent=2)\n        print(f\"Training history saved to {save_path}\")\n    \n    def load_history(self, load_path: str):\n        \"\"\"Load training history\"\"\"\n        with open(load_path, 'r') as f:\n            history = json.load(f)\n        \n        self.train_history = defaultdict(list, history['train'])\n        self.val_history = defaultdict(list, history['val'])\n        self.test_history = defaultdict(list, history.get('test', {}))\n        print(f\"Training history loaded from {load_path}\")\n    \n    def generate_report(self, save_path: Optional[str] = None) -> str:\n        \"\"\"Generate a comprehensive classification report\"\"\"\n        if len(self.predictions) == 0:\n            return \"No predictions available for report generation\"\n        \n        predictions = np.array(self.predictions)\n        targets = np.array(self.targets)\n        \n        # Get current metrics\n        metrics = self.compute_metrics()\n        \n        # Generate classification report\n        class_report = classification_report(targets, predictions, zero_division=0)\n        \n        # Create comprehensive report\n        report = f\"\"\"\n=== CLASSIFICATION REPORT ===\n\nOverall Metrics:\n{metrics}\n\nDetailed Classification Report:\n{class_report}\n\nConfusion Matrix:\n{confusion_matrix(targets, predictions)}\n\"\"\"\n        \n        if save_path:\n            with open(save_path, 'w') as f:\n                f.write(report)\n            print(f\"Report saved to {save_path}\")\n        \n        return report\n    \n    def create_summary_plots(self, output_dir: str):\n        \"\"\"Create all summary plots and save them\"\"\"\n        output_dir = Path(output_dir)\n        output_dir.mkdir(parents=True, exist_ok=True)\n        \n        # Training curves\n        if len(self.train_history['loss']) > 0 or len(self.val_history['loss']) > 0:\n            self.plot_training_curves(\n                save_path=output_dir / 'training_curves.png',\n                show=False\n            )\n        \n        # Confusion matrix\n        if len(self.predictions) > 0:\n            self.plot_confusion_matrix(\n                save_path=output_dir / 'confusion_matrix.png',\n                show=False\n            )\n        \n        # ROC curve\n        if self.task_type == 'binary' and len(self.predictions) > 0:\n            self.plot_roc_curve(\n                save_path=output_dir / 'roc_curve.png',\n                show=False\n            )\n            \n            self.plot_precision_recall_curve(\n                save_path=output_dir / 'pr_curve.png',\n                show=False\n            )\n        \n        print(f\"All plots saved to {output_dir}\")\n\n\nclass EarlyStopping:\n    \"\"\"Early stopping utility class\"\"\"\n    \n    def __init__(\n        self,\n        patience: int = 10,\n        min_delta: float = 1e-4,\n        metric: str = 'auc_roc',\n        mode: str = 'max',\n        restore_best_weights: bool = True\n    ):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.metric = metric\n        self.mode = mode\n        self.restore_best_weights = restore_best_weights\n        \n        self.best_score = None\n        self.best_epoch = 0\n        self.wait = 0\n        self.stopped_epoch = 0\n        self.best_weights = None\n        \n        if mode == 'max':\n            self.monitor_op = np.greater\n            self.min_delta *= 1\n        else:\n            self.monitor_op = np.less\n            self.min_delta *= -1\n    \n    def __call__(self, current_score: float, model_state_dict: dict = None, epoch: int = 0) -> bool:\n        \"\"\"\n        Check if training should be stopped\n        \n        Returns:\n            True if training should be stopped, False otherwise\n        \"\"\"\n        if self.best_score is None:\n            self.best_score = current_score\n            self.best_epoch = epoch\n            if model_state_dict is not None:\n                self.best_weights = model_state_dict.copy()\n            return False\n        \n        if self.monitor_op(current_score - self.min_delta, self.best_score):\n            self.best_score = current_score\n            self.best_epoch = epoch\n            self.wait = 0\n            if model_state_dict is not None:\n                self.best_weights = model_state_dict.copy()\n        else:\n            self.wait += 1\n            \n        if self.wait >= self.patience:\n            self.stopped_epoch = epoch\n            return True\n        \n        return False\n    \n    def get_best_weights(self):\n        \"\"\"Get the best model weights\"\"\"\n        return self.best_weights\n\n","metadata":{"trusted":true,"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true,"editable":false},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## config","metadata":{"editable":false}},{"cell_type":"code","source":"import pandas as pd\n\ndf = pd.read_csv(\"/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv\")\ntest_path = \"/kaggle/input/rsna-intracranial-aneurysm-detection/series/1.2.826.0.1.3680043.8.498.10004044428023505108375152878107656647\"\nvolum, meta = process_dicom_series(test_path, df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T15:04:15.852732Z","iopub.execute_input":"2025-08-27T15:04:15.853076Z","iopub.status.idle":"2025-08-27T15:04:19.397643Z","shell.execute_reply.started":"2025-08-27T15:04:15.853053Z","shell.execute_reply":"2025-08-27T15:04:19.396719Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"meta, volum.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T15:04:35.149491Z","iopub.execute_input":"2025-08-27T15:04:35.150314Z","iopub.status.idle":"2025-08-27T15:04:35.15605Z","shell.execute_reply.started":"2025-08-27T15:04:35.150287Z","shell.execute_reply":"2025-08-27T15:04:35.154849Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df.iloc[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T15:15:48.391866Z","iopub.execute_input":"2025-08-27T15:15:48.392288Z","iopub.status.idle":"2025-08-27T15:15:48.400254Z","shell.execute_reply.started":"2025-08-27T15:15:48.39226Z","shell.execute_reply":"2025-08-27T15:15:48.399363Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nplot_image_grid(\n    volum,\n    (5, 5)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-08-27T15:14:14.874731Z","iopub.execute_input":"2025-08-27T15:14:14.875075Z","iopub.status.idle":"2025-08-27T15:14:16.066886Z","shell.execute_reply.started":"2025-08-27T15:14:14.875054Z","shell.execute_reply":"2025-08-27T15:14:16.065977Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Hello\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-01T10:06:01.40985Z","iopub.execute_input":"2025-09-01T10:06:01.410142Z","iopub.status.idle":"2025-09-01T10:06:01.417193Z","shell.execute_reply.started":"2025-09-01T10:06:01.410107Z","shell.execute_reply":"2025-09-01T10:06:01.416509Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Start Here","metadata":{"editable":false}},{"cell_type":"code","source":"!git clone git@github.com:mohame54/test_rnsa.git","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-01T10:09:05.820422Z","iopub.execute_input":"2025-09-01T10:09:05.82069Z","iopub.status.idle":"2025-09-01T10:09:37.967053Z","shell.execute_reply.started":"2025-09-01T10:09:05.820658Z","shell.execute_reply":"2025-09-01T10:09:37.96639Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nos.listdir(\".\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-01T10:09:59.765616Z","iopub.execute_input":"2025-09-01T10:09:59.766254Z","iopub.status.idle":"2025-09-01T10:09:59.771764Z","shell.execute_reply.started":"2025-09-01T10:09:59.766224Z","shell.execute_reply":"2025-09-01T10:09:59.771275Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ID_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# All tags (other than PixelData and SeriesInstanceUID) that may be in a test set dcm file\nDICOM_TAG_ALLOWLIST = [\n    'BitsAllocated',\n    'BitsStored',\n    'Columns',\n    'FrameOfReferenceUID',\n    'HighBit',\n    'ImageOrientationPatient',\n    'ImagePositionPatient',\n    'InstanceNumber',\n    'Modality',\n    'PatientID',\n    'PhotometricInterpretation',\n    'PixelRepresentation',\n    'PixelSpacing',\n    'PlanarConfiguration',\n    'RescaleIntercept',\n    'RescaleSlope',\n    'RescaleType',\n    'Rows',\n    'SOPClassUID',\n    'SOPInstanceUID',\n    'SamplesPerPixel',\n    'SliceThickness',\n    'SpacingBetweenSlices',\n    'StudyInstanceUID',\n    'TransferSyntaxUID',\n]\n\n# Replace this function with your inference code.\n# You can return either a Pandas or Polars dataframe, though Polars is recommended.\n# Each prediction (except the very first) must be returned within 30 minutes of the series being provided.\ndef predict(series_path: str) -> pl.DataFrame | pd.DataFrame:\n    \"\"\"Make a prediction.\"\"\"\n    # --------- Replace this section with your own prediction code ---------\n    series_id = os.path.basename(series_path)\n    \n    all_filepaths = []\n    for root, _, files in os.walk(series_path):\n        for file in files:\n            if file.endswith('.dcm'):\n                all_filepaths.append(os.path.join(root, file))\n    all_filepaths.sort()\n    \n    # Collect tags from the dicoms\n    tags = defaultdict(list)\n    tags['SeriesInstanceUID'] = series_id\n    global dcms\n    for filepath in all_filepaths:\n        ds = pydicom.dcmread(filepath, force=True)\n        tags['filepath'].append(filepath)\n        for tag in DICOM_TAG_ALLOWLIST:\n            tags[tag].append(getattr(ds, tag, None))\n        # The image is in ds.PixelData\n\n    # ... do some machine learning magic ...\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\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 ------------------------------\n    # You MUST have the following code in your `predict` function\n    # to prevent \"out of disk space\" errors. This is a temporary workaround\n    # as we implement improvements to our evaluation system.\n    shutil.rmtree('/kaggle/shared', ignore_errors=True)\n    # ----------------------------------------------------------------------\n    \n    return predictions.drop(ID_COL)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-28T20:30:17.20469Z","iopub.execute_input":"2025-07-28T20:30:17.205112Z","iopub.status.idle":"2025-07-28T20:30:17.236447Z","shell.execute_reply.started":"2025-07-28T20:30:17.205079Z","shell.execute_reply":"2025-07-28T20:30:17.23539Z"},"editable":false},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"When your notebook is run on the hidden test set, `inference_server.serve` must be called within 15 minutes of the notebook starting or the gateway will throw an error. If you need more than 15 minutes to load your model you can do so during the very first `predict` call.","metadata":{"editable":false}},{"cell_type":"code","source":"inference_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":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-28T20:30:17.237419Z","iopub.execute_input":"2025-07-28T20:30:17.237805Z","iopub.status.idle":"2025-07-28T20:30:32.279895Z","shell.execute_reply.started":"2025-07-28T20:30:17.237775Z","shell.execute_reply":"2025-07-28T20:30:32.278943Z"},"editable":false},"outputs":[],"execution_count":null}]}