{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaL4","dataSources":[{"sourceId":91496,"databundleVersionId":11802066,"sourceType":"competition"},{"sourceId":99552,"databundleVersionId":13851420,"sourceType":"competition"},{"sourceId":13267216,"sourceType":"datasetVersion","datasetId":8343284}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!uv pip install monai dynamic_network_architectures cucim-cu12 fire","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:15.630439Z","iopub.execute_input":"2025-10-05T12:58:15.630906Z","iopub.status.idle":"2025-10-05T12:58:30.050817Z","shell.execute_reply.started":"2025-10-05T12:58:15.630883Z","shell.execute_reply":"2025-10-05T12:58:30.050178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!mkdir inference","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:30.051913Z","iopub.execute_input":"2025-10-05T12:58:30.05212Z","iopub.status.idle":"2025-10-05T12:58:30.165632Z","shell.execute_reply.started":"2025-10-05T12:58:30.0521Z","shell.execute_reply":"2025-10-05T12:58:30.164992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile inference/_gpu_resampling.py\n\nfrom copy import deepcopy\nfrom typing import Union, Tuple, List\n\nimport numpy as np\nimport cupy as cp\nfrom cucim.skimage.transform import resize as cucim_resize\nfrom cupyx.scipy.ndimage import map_coordinates as cucim_map_coordinates\n\nANISO_THRESHOLD = 3\n\n\ndef get_do_separate_z(\n    spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n    anisotropy_threshold=ANISO_THRESHOLD,\n):\n    do_separate_z = (np.max(spacing) / np.min(spacing)) > anisotropy_threshold\n    return do_separate_z\n\n\ndef get_lowres_axis(new_spacing: Union[Tuple[float, ...], List[float], np.ndarray]):\n    axis = np.where(max(new_spacing) / np.array(new_spacing) == 1)[0]\n    return axis\n\n\ndef compute_new_shape(\n    old_shape: Union[Tuple[int, ...], List[int], np.ndarray],\n    old_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n    new_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n) -> np.ndarray:\n    assert len(old_spacing) == len(old_shape)\n    assert len(old_shape) == len(new_spacing)\n    new_shape = np.array(\n        [int(round(i / j * k)) for i, j, k in zip(old_spacing, new_spacing, old_shape)]\n    )\n    return new_shape\n\n\ndef determine_do_sep_z_and_axis(\n    force_separate_z: bool,\n    current_spacing,\n    new_spacing,\n    separate_z_anisotropy_threshold: float = ANISO_THRESHOLD,\n) -> Tuple[bool, Union[int, None]]:\n    if force_separate_z is not None:\n        do_separate_z = force_separate_z\n        if force_separate_z:\n            axis = get_lowres_axis(current_spacing)\n        else:\n            axis = None\n    else:\n        if get_do_separate_z(current_spacing, separate_z_anisotropy_threshold):\n            do_separate_z = True\n            axis = get_lowres_axis(current_spacing)\n        elif get_do_separate_z(new_spacing, separate_z_anisotropy_threshold):\n            do_separate_z = True\n            axis = get_lowres_axis(new_spacing)\n        else:\n            do_separate_z = False\n            axis = None\n\n    if axis is not None:\n        if len(axis) == 3:\n            do_separate_z = False\n            axis = None\n        elif len(axis) == 2:\n            do_separate_z = False\n            axis = None\n        else:\n            axis = axis[0]\n    return do_separate_z, axis\n\n\ndef resample_data_or_seg_to_spacing(\n    data: np.ndarray,\n    current_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n    new_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n    order: int = 3,\n    order_z: int = 0,\n    force_separate_z: Union[bool, None] = False,\n    separate_z_anisotropy_threshold: float = ANISO_THRESHOLD,\n):\n    do_separate_z, axis = determine_do_sep_z_and_axis(\n        force_separate_z, current_spacing, new_spacing, separate_z_anisotropy_threshold\n    )\n\n    if data is not None:\n        assert data.ndim == 4, \"data must be c x y z\"\n\n    shape = np.array(data.shape)\n    new_shape = compute_new_shape(shape[1:], current_spacing, new_spacing)\n\n    data_reshaped = resample_data_or_seg(\n        data, new_shape, axis, order, do_separate_z, order_z=order_z\n    )\n    return data_reshaped\n\n\ndef resample_data_or_seg_to_shape(\n    data: np.ndarray,\n    new_shape: Union[Tuple[int, ...], List[int], np.ndarray],\n    current_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n    new_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n    order: int = 3,\n    order_z: int = 0,\n    force_separate_z: Union[bool, None] = False,\n    separate_z_anisotropy_threshold: float = ANISO_THRESHOLD,\n):\n    \"\"\"\n    needed for segmentation export. Stupid, I know\n    \"\"\"\n    do_separate_z, axis = determine_do_sep_z_and_axis(\n        force_separate_z, current_spacing, new_spacing, separate_z_anisotropy_threshold\n    )\n\n    if data is not None:\n        assert data.ndim == 4, \"data must be c x y z\"\n\n    data_reshaped = resample_data_or_seg(\n        data, new_shape, axis, order, do_separate_z, order_z=order_z\n    )\n    return data_reshaped\n\n\ndef resample_data_or_seg(\n    data: np.ndarray,\n    _new_shape: Union[Tuple[float, ...], List[float], np.ndarray],\n    axis: Union[None, int] = None,\n    order: int = 3,\n    do_separate_z: bool = False,\n    order_z: int = 0,\n    dtype_out=None,\n):\n    \"\"\"\n    cuCIM/cupy-accelerated version of resample_data_or_seg\n    separate_z=True will resample with order 0 along z\n    :param data: numpy array (c, x, y, z)\n    :param new_shape:\n    :param axis:\n    :param order:\n    :param do_separate_z:\n    :param order_z: only applies if do_separate_z is True\n    :return: numpy array\n    \"\"\"\n    assert data.ndim == 4, \"data must be (c, x, y, z)\"\n    assert len(_new_shape) == data.ndim - 1\n\n    # Convert to GPU\n    data_gpu = cp.asarray(data, dtype=cp.float32)\n    kwargs = {\"mode\": \"edge\", \"anti_aliasing\": False}\n    shape = cp.array(data_gpu[0].shape)\n    new_shape = cp.array(_new_shape)\n\n    if dtype_out is None:\n        dtype_out = cp.float32\n\n    reshaped_final = cp.zeros(\n        tuple([data_gpu.shape[0]] + new_shape.tolist()), dtype=dtype_out\n    )\n\n    if cp.any(shape != new_shape):\n        data_gpu = data_gpu.astype(cp.float32, copy=False)\n\n        if do_separate_z:\n            assert (\n                axis is not None\n            ), \"If do_separate_z, we need to know what axis is anisotropic\"\n\n            if axis == 0:\n                new_shape_2d = new_shape[1:]\n            elif axis == 1:\n                new_shape_2d = new_shape[[0, 2]]\n            else:\n                new_shape_2d = new_shape[:-1]\n\n            for c in range(data_gpu.shape[0]):\n                tmp = deepcopy(new_shape)\n                tmp[axis] = shape[axis]\n                reshaped_here = cp.zeros(tuple(tmp.tolist()))\n\n                # GPU-accelerated slice processing with cuCIM\n                for slice_id in range(int(shape[axis])):\n                    if axis == 0:\n                        reshaped_here[slice_id] = cucim_resize(\n                            data_gpu[c, slice_id], new_shape_2d, order, **kwargs\n                        )\n                    elif axis == 1:\n                        reshaped_here[:, slice_id] = cucim_resize(\n                            data_gpu[c, :, slice_id], new_shape_2d, order, **kwargs\n                        )\n                    else:\n                        reshaped_here[:, :, slice_id] = cucim_resize(\n                            data_gpu[c, :, :, slice_id], new_shape_2d, order, **kwargs\n                        )\n\n                if shape[axis] != new_shape[axis]:\n                    # GPU-accelerated coordinate mapping with cupyx\n                    rows, cols, dim = new_shape[0], new_shape[1], new_shape[2]\n                    orig_rows, orig_cols, orig_dim = reshaped_here.shape\n\n                    # align_corners=False - same logic as original\n                    row_scale = float(orig_rows) / rows\n                    col_scale = float(orig_cols) / cols\n                    dim_scale = float(orig_dim) / dim\n\n                    map_rows, map_cols, map_dims = cp.mgrid[:rows, :cols, :dim]\n                    map_rows = row_scale * (map_rows + 0.5) - 0.5\n                    map_cols = col_scale * (map_cols + 0.5) - 0.5\n                    map_dims = dim_scale * (map_dims + 0.5) - 0.5\n\n                    coord_map = cp.array([map_rows, map_cols, map_dims])\n\n                    # GPU-accelerated coordinate mapping\n                    reshaped_final[c] = cucim_map_coordinates(\n                        reshaped_here, coord_map, order=order_z, mode=\"nearest\"\n                    )[None]\n                else:\n                    reshaped_final[c] = reshaped_here\n        else:\n            # GPU-accelerated direct resize\n            for c in range(data_gpu.shape[0]):\n                reshaped_final[c] = cucim_resize(\n                    data_gpu[c], new_shape, order, **kwargs\n                )\n\n        # Convert back to CPU numpy array\n        return cp.asnumpy(reshaped_final)\n    else:\n        # No resampling needed\n        return data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:30.166438Z","iopub.execute_input":"2025-10-05T12:58:30.16681Z","iopub.status.idle":"2025-10-05T12:58:30.174606Z","shell.execute_reply.started":"2025-10-05T12:58:30.166774Z","shell.execute_reply":"2025-10-05T12:58:30.174087Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile inference/_gpu_seg_resampling.py\n\nfrom typing import Union, Tuple, List\nimport numpy as np\nimport cupy as cp\nimport torch\nfrom cucim.skimage.transform import resize as cucim_resize\n\n\ndef compute_new_shape(\n    old_shape: Union[Tuple[int, ...], List[int], np.ndarray],\n    old_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n    new_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n) -> np.ndarray:\n    \"\"\"Compute new shape after resampling\"\"\"\n    assert len(old_spacing) == len(old_shape)\n    assert len(old_shape) == len(new_spacing)\n    new_shape = np.array(\n        [int(round(i / j * k)) for i, j, k in zip(old_spacing, new_spacing, old_shape)]\n    )\n    return new_shape\n\n\ndef resample_segmentation_to_spacing(\n    data: np.ndarray,\n    current_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n    new_spacing: Union[Tuple[float, ...], List[float], np.ndarray],\n):\n    \"\"\"\n    Resample segmentation to new spacing using nearest neighbor\n\n    Args:\n        data: segmentation data with shape (c, x, y, z)\n        current_spacing: current spacing (x, y, z)\n        new_spacing: target spacing (x, y, z)\n\n    Returns:\n        resampled segmentation with same channel dimension\n    \"\"\"\n    if data is not None:\n        assert data.ndim == 4, \"data must be c x y z\"\n\n    shape = np.array(data.shape)\n    new_shape = compute_new_shape(shape[1:], current_spacing, new_spacing)\n\n    return resample_segmentation_to_shape(data, new_shape)\n\n\ndef resize_segmentation_cupy(segmentation, new_shape, order=3):\n    \"\"\"\n    CuPy version of resize_segmentation\n    Input: CuPy array\n    Output: CuPy array\n    \"\"\"\n    tpe = segmentation.dtype\n    assert len(segmentation.shape) == len(\n        new_shape\n    ), \"new shape must have same dimensionality as segmentation\"\n\n    if order == 0:\n        return cucim_resize(\n            segmentation.astype(cp.float32),\n            new_shape,\n            order,\n            mode=\"edge\",\n            clip=True,\n            anti_aliasing=False,\n        ).astype(tpe)\n    else:\n        reshaped = cp.zeros(new_shape, dtype=segmentation.dtype)\n\n        unique_labels = cp.sort(cp.unique(segmentation.ravel()))\n        for i, c in enumerate(unique_labels):\n            mask = segmentation == c\n            reshaped_multihot = cucim_resize(\n                mask.astype(cp.float32),\n                new_shape,\n                order,\n                mode=\"edge\",\n                clip=True,\n                anti_aliasing=False,\n            )\n            reshaped[reshaped_multihot >= 0.5] = c\n        return reshaped\n\n\ndef resample_segmentation_to_shape(\n    data: Union[torch.Tensor, np.ndarray],\n    new_shape: Union[Tuple[int, ...], List[int], np.ndarray],\n):\n    \"\"\"\n    Resample segmentation to new shape using nearest neighbor\n\n    Args:\n        data: segmentation data with shape (c, x, y, z)\n        new_shape: target shape (x, y, z)\n\n    Returns:\n        resampled segmentation\n    \"\"\"\n    if isinstance(data, torch.Tensor):\n        data = data.numpy()\n\n    if data is not None:\n        assert data.ndim == 4, \"data must be c x y z\"\n\n    return _resample_segmentation_core_cupy(data, new_shape)\n\n\ndef _resample_segmentation_core_cupy(\n    data: np.ndarray,\n    new_shape: Union[Tuple[int, ...], List[int], np.ndarray],\n    dtype_out=None,\n):\n    \"\"\"\n    Core segmentation resampling function - CuPy accelerated\n    \"\"\"\n    assert data.ndim == 4, \"data must be (c, x, y, z)\"\n    assert len(new_shape) == data.ndim - 1\n\n    shape = np.array(data[0].shape)\n    new_shape_np = np.array(new_shape)\n\n    if dtype_out is None:\n        dtype_out = data.dtype\n\n    # Early return if no resampling needed\n    if np.all(shape == new_shape_np):\n        return data\n\n    # Convert to GPU\n    data_gpu = cp.asarray(data, dtype=cp.float32)\n\n    # Allocate output on GPU\n    reshaped_final_gpu = cp.zeros((data_gpu.shape[0], *new_shape), dtype=cp.float32)\n\n    # Direct resize with nearest neighbor for all channels\n    for c in range(data_gpu.shape[0]):\n        reshaped_final_gpu[c] = resize_segmentation_cupy(\n            data_gpu[c], new_shape, order=0  # Always nearest neighbor for segmentation\n        )\n\n    return cp.asnumpy(reshaped_final_gpu).astype(dtype_out)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:30.175842Z","iopub.execute_input":"2025-10-05T12:58:30.176017Z","iopub.status.idle":"2025-10-05T12:58:30.187199Z","shell.execute_reply.started":"2025-10-05T12:58:30.176003Z","shell.execute_reply":"2025-10-05T12:58:30.186705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile inference/_transform.py\n\nimport numpy as np\nimport torch\nfrom typing import Union, Optional\nfrom monai.transforms import Transform\nfrom monai.utils import convert_to_tensor\nfrom skimage.morphology import remove_small_objects\n\n\nclass RemoveSmallObjectsDense(Transform):\n    \"\"\"\n    Remove small connected components from dense label image.\n    Works directly on labeled arrays without one-hot conversion.\n\n    Args:\n        min_size: Minimum component size in voxels to keep.\n        connectivity: Neighborhood connectivity (1, 2, or 3 for 3D images).\n        device: 'cpu' or 'cuda'. Auto-detects from input if None.\n    \"\"\"\n\n    def __init__(\n        self, min_size: int = 64, connectivity: int = 1, device: Optional[str] = None\n    ):\n        super().__init__()\n        self.min_size = min_size\n        self.connectivity = connectivity\n        self.device = device\n\n    def __call__(self, img: Union[np.ndarray, torch.Tensor]) -> torch.Tensor:\n        \"\"\"\n        Args:\n            img: Dense label image of shape (H, W, D) with integer labels.\n\n        Returns:\n            Image with small objects removed.\n        \"\"\"\n        device = self.device or (\n            \"cuda\" if isinstance(img, torch.Tensor) and img.is_cuda else \"cpu\"\n        )\n\n        if device == \"cuda\":\n            return self._process_cpu(\n                img.to(device=\"cpu\") if isinstance(img, torch.Tensor) else img\n            )\n        return self._process_cpu(img)\n\n    def _process_cpu(self, img: Union[np.ndarray, torch.Tensor]) -> torch.Tensor:\n        \"\"\"CPU implementation using scikit-image.\"\"\"\n        img_np = img.cpu().numpy() if isinstance(img, torch.Tensor) else img\n\n        # scikit-image's remove_small_objects works directly on labeled images\n        cleaned = img_np * remove_small_objects(\n            img_np.astype(bool), min_size=self.min_size, connectivity=self.connectivity\n        )\n\n        return convert_to_tensor(cleaned, dtype=torch.uint8)\n\n    def _process_gpu(self, img: Union[np.ndarray, torch.Tensor]) -> torch.Tensor:\n        raise NotImplementedError(\"GPU version not implemented yet.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:30.187667Z","iopub.execute_input":"2025-10-05T12:58:30.187832Z","iopub.status.idle":"2025-10-05T12:58:30.198042Z","shell.execute_reply.started":"2025-10-05T12:58:30.187819Z","shell.execute_reply":"2025-10-05T12:58:30.197547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile inference/inferer.py\n\nimport time\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom typing import List, cast\nfrom pathlib import Path\nfrom monai.transforms import NormalizeIntensity\nfrom monai.inferers.utils import sliding_window_inference\nfrom dynamic_network_architectures.architectures.unet import (\n    ResidualEncoderUNet,\n    UNetDecoder,\n)\nimport nibabel as nib\nimport monai.transforms as mt\nfrom inference._gpu_resampling import (\n    resample_data_or_seg,\n    compute_new_shape,\n    determine_do_sep_z_and_axis,\n)\nfrom inference._transform import RemoveSmallObjectsDense\n\nfrom inference._gpu_seg_resampling import resample_segmentation_to_spacing\n\n\ndef add_additional_channel_to_state_dict(model_weight, model):\n    state_dict = model_weight[\"network_weights\"]\n\n    # Get current model's state dict to identify which keys need modification\n    model_state = model.state_dict()\n\n    # Find all layers with channel dimension mismatch\n    for key in list(state_dict.keys()):\n        if key in model_state:\n            pretrained_shape = state_dict[key].shape\n            current_shape = model_state[key].shape\n\n            # Check if this is a weight tensor with channel dimension mismatch\n            if len(pretrained_shape) == 5 and pretrained_shape != current_shape:\n                # This is likely a conv weight: [out_channels, in_channels, D, H, W]\n                if pretrained_shape[1] == 1 and current_shape[1] == 2:\n                    print(f\"Expanding {key} from {pretrained_shape} to {current_shape}\")\n\n                    # Create new tensor with correct shape\n                    new_weight = torch.zeros(current_shape)\n\n                    # Copy pretrained weights to channel 0\n                    new_weight[:, 0:1, :, :, :] = state_dict[key]\n\n                    # Initialize channel 1 with N(0, 0.01)\n                    torch.nn.init.normal_(\n                        new_weight[:, 1:2, :, :, :], mean=0.0, std=0.01\n                    )\n\n                    # Replace in state dict\n                    state_dict[key] = new_weight\n\n    return state_dict\n\n\nclass ResEncoderUNetModel(nn.Module):\n    def __init__(\n        self,\n        input_channels: int,\n        num_classes: int,\n        deep_supervision: bool = True,\n        pretrained_weights_path: str | None = None,\n        keep_decoder_weights: bool = True,\n    ):\n        super().__init__()\n        self.deep_supervision = deep_supervision\n\n        self.model = ResidualEncoderUNet(\n            input_channels=input_channels,\n            n_stages=5,\n            num_classes=num_classes,\n            features_per_stage=[32, 64, 128, 256, 320],\n            conv_op=nn.Conv3d,\n            kernel_sizes=(3, 3, 3, 3, 3),\n            strides=((1, 1, 1), (2, 2, 2), (2, 2, 2), (2, 2, 2), (1, 2, 2)),  # type: ignore\n            n_blocks_per_stage=(1, 3, 4, 6, 6),\n            n_conv_per_stage_decoder=(1, 1, 1, 1),\n            conv_bias=True,\n            norm_op=nn.InstanceNorm3d,\n            norm_op_kwargs={\"affine\": True, \"eps\": 1e-3},\n            dropout_op=nn.Dropout3d,\n            dropout_op_kwargs={\"p\": 0.20},\n            nonlin=nn.LeakyReLU,\n            nonlin_kwargs={\"inplace\": True},\n            deep_supervision=deep_supervision,\n        )\n\n        if pretrained_weights_path is not None:\n            model_weight = torch.load(pretrained_weights_path, weights_only=False)\n            for key in list(model_weight[\"network_weights\"].keys()):\n                if key.startswith(\"decoder.seg_layer\"):\n                    del model_weight[\"network_weights\"][key]\n\n            state_dict = add_additional_channel_to_state_dict(model_weight, self.model)\n\n            self.model.load_state_dict(state_dict, strict=False)\n\n            if not keep_decoder_weights:\n                print(\"Reinitializing decoder weights...\")\n                self.model.decoder = UNetDecoder(\n                    encoder=self.model.encoder,\n                    num_classes=num_classes,\n                    n_conv_per_stage=(1, 1, 1, 1),\n                    deep_supervision=deep_supervision,\n                )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor | List[torch.Tensor]:\n\n        output = self.model(x)\n        if self.deep_supervision:\n            if self.training:\n                return output[::-1]\n            else:\n                return output[0]\n\n        return output\n\n\nclass Inferer:\n    def __init__(\n        self,\n        model_save_path: str,\n        device=\"cuda\",\n        inference_roi_sizes: list[tuple[int, int, int]] = [(40, 160, 160)],\n        inference_batch_size: int = 8,\n        inference_overlap=0.25,\n        inference_dtype=torch.float32,\n        use_amp: bool = False,\n    ):\n        self.model_save_path = model_save_path\n        self.device = torch.device(device)\n        self.inference_roi_sizes = inference_roi_sizes\n        self.inference_batch_size = inference_batch_size\n        self.inference_overlap = inference_overlap\n        self.inference_dtype = inference_dtype\n        self.use_amp = use_amp\n\n        if inference_dtype == torch.float32:\n            print(\n                \"[WARNING] Cannot use mixed precision with float32, setting use_amp to False.\"\n            )\n            self.use_amp = False\n\n        self._post_transforms = mt.Compose(\n            [\n                mt.AsDiscrete(argmax=True),\n                mt.EnsureType(dtype=torch.uint8),\n                RemoveSmallObjectsDense(min_size=500, connectivity=5),\n            ],\n        )\n        self._model = self._load_model()\n\n    def _load_model(self):\n\n        model = ResEncoderUNetModel(\n            input_channels=1,\n            num_classes=15,\n            deep_supervision=False,\n            pretrained_weights_path=None,\n        )\n\n        weights = torch.load(self.model_save_path, map_location=self.device)\n        model.load_state_dict(weights[\"network_weights\"])\n        model.to(self.device)\n        model.eval()\n\n        return model\n\n    def infer(self, volume: torch.Tensor) -> torch.Tensor:\n\n        volume = volume.unsqueeze(0).to(device=self.device)  # (1, C, H, W, D)\n        output_sum = None\n        num_scales = len(self.inference_roi_sizes)\n    \n        for inference_size in self.inference_roi_sizes:\n            if self.use_amp:\n                with torch.autocast(\n                    device_type=self.device.type, dtype=self.inference_dtype\n                ):\n                    with torch.inference_mode():\n                        output = sliding_window_inference(\n                            inputs=volume,\n                            roi_size=inference_size,\n                            sw_batch_size=self.inference_batch_size,\n                            predictor=lambda x: self._model(x)[:, :-1, ...], # remove auxiliary head\n                            overlap=self.inference_overlap,\n                            mode=\"gaussian\",\n                            sw_device=self.device,\n                            # device=self.device,\n                            device=\"cpu\",\n                        )\n    \n            else:\n                with torch.inference_mode():\n                    output = sliding_window_inference(\n                        inputs=volume,\n                        roi_size=inference_size,\n                        sw_batch_size=self.inference_batch_size,\n                        predictor=self._model,\n                        overlap=self.inference_overlap,\n                        mode=\"gaussian\",\n                        sw_device=self.device,\n                        # device=self.device,\n                        device=\"cpu\"\n                    )\n    \n            assert isinstance(output, torch.Tensor)\n            output = output.squeeze(0)\n            \n            # Accumulate in-place instead of appending to list\n            if output_sum is None:\n                output_sum = output.clone()  # First iteration: create accumulator\n            else:\n                output_sum.add_(output)  # Subsequent iterations: in-place addition\n            \n        # Compute average in-place\n        output = output_sum.div_(num_scales)\n        output = self._post_transforms(output)\n    \n        return output.detach().cpu()\n\n\ndef _reorient_to_lps(volume, info_dict={}):\n    \"\"\"Reorient volume and segmentation to LPS coordinate system.\"\"\"\n    orientation = nib.orientations.io_orientation(volume.affine)\n    lps_orientation = nib.orientations.axcodes2ornt((\"L\", \"P\", \"S\"))\n    transform = nib.orientations.ornt_transform(orientation, lps_orientation)\n\n    info_dict[\"original_orientation\"] = nib.orientations.ornt2axcodes(orientation)\n    info_dict[\"target_orientation\"] = \"LPS\"\n\n    volume_reoriented = nib.orientations.apply_orientation(\n        volume.get_fdata(), transform\n    )\n\n    volume_affine = np.dot(\n        volume.affine, nib.orientations.inv_ornt_aff(transform, volume.shape)\n    )\n\n    return nib.Nifti1Image(volume_reoriented, volume_affine)\n\n\ndef _resample_volume_nnunet(volume, spacing=(1.0, 1.0, 1.0), info_dict={}):\n    \"\"\"\n    CHANGE: Replaced simple scipy zoom with nnUNet's anisotropic-aware resampling\n    \"\"\"\n    current_spacing = volume.header.get_zooms()[:3]\n    # Get data and add channel dimension for nnUNet format (c, x, y, z)\n    volume_data = volume.get_fdata()[np.newaxis, ...]  # Add channel dim\n\n    # Compute new shape using nnUNet utilities\n    new_shape = compute_new_shape(volume_data.shape[1:], current_spacing, spacing)\n\n    # Determine if separate z-axis handling is needed\n    do_separate_z, axis = determine_do_sep_z_and_axis(None, current_spacing, spacing)\n\n    volume_resampled = resample_data_or_seg(\n        volume_data,\n        new_shape,\n        axis=axis,\n        order=3,\n        do_separate_z=do_separate_z,\n        order_z=0,\n    )\n\n    # Remove channel dimension\n    volume_resampled = volume_resampled[0]\n\n    # FIXED: Proper affine matrix handling\n    zoom_factors = np.array(current_spacing) / np.array(spacing)\n\n    volume_affine = volume.affine.copy()\n\n    # Scale the existing affine matrix to account for resampling\n    # This preserves rotation, shear, and orientation information\n    volume_affine[:3, :3] = volume_affine[:3, :3] / zoom_factors\n    volume_img = nib.Nifti1Image(volume_resampled, volume_affine)\n\n    return volume_img\n\n\ndef _get_preprocessed_volume(\n    volume_path: str,\n) -> nib.Nifti1Image:\n    nii_volume = nib.load(volume_path)\n    nii_npy = nii_volume.get_fdata()\n    if len(nii_npy.shape) == 4:\n        # Take first channel if multiple channels exist\n        nii_npy = nii_npy[..., 0]\n        nii_volume = nib.Nifti1Image(nii_npy, nii_volume.affine, nii_volume.header)\n    nii_volume = _reorient_to_lps(nii_volume)\n    nii_volume = _resample_volume_nnunet(nii_volume, spacing=(0.4425, 0.4425, 0.80))\n    return nii_volume\n\n\ndef generate_segmentation(\n    volume: np.ndarray,\n    inferer: Inferer,\n    current_spacing: tuple[float, float, float],\n    target_spacing: tuple[float, float, float],\n) -> np.ndarray:\n\n    volume_data = torch.from_numpy(volume)\n    volume_data = volume_data.unsqueeze(0)\n    volume_data = cast(torch.Tensor, NormalizeIntensity()(volume_data))\n    volume_data = volume_data.permute(0, 3, 1, 2)  # (C, D, H, W)\n    start_time = time.time()\n    segmentation = inferer.infer(volume_data)\n    print(f\"Finished generation in {time.time() - start_time:.2f} seconds.\")\n    segmentation = segmentation.numpy()\n    segmentation = np.transpose(segmentation, (0, 2, 3, 1))\n\n    segmentation = resample_segmentation_to_spacing(\n        segmentation,\n        current_spacing=current_spacing,\n        new_spacing=target_spacing,\n    )\n    segmentation = segmentation.astype(np.uint8)\n    return segmentation\n\n\ndef generate_and_write_segmentation(\n    volume_path: str,\n    inferer: Inferer,\n    output_path: str,\n    target_spacing: tuple[float, float, float],\n):\n    nii_volume = _get_preprocessed_volume(volume_path)\n    assert nii_volume.affine is not None\n\n    current_spacing = nii_volume.header.get_zooms()[:3]\n    assert len(current_spacing) == 3\n    segmentation = generate_segmentation(\n        nii_volume.get_fdata(),\n        inferer,\n        current_spacing=current_spacing,\n        target_spacing=target_spacing,\n    )\n    segmentation = segmentation[0]\n\n    zoom_factors = np.array(current_spacing) / np.array(target_spacing)\n    seg_affine = nii_volume.affine.copy()\n    seg_affine[:3, :3] = seg_affine[:3, :3] / zoom_factors\n    seg_img = nib.Nifti1Image(segmentation, seg_affine)\n    nib.save(seg_img, output_path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:30.198591Z","iopub.execute_input":"2025-10-05T12:58:30.198746Z","iopub.status.idle":"2025-10-05T12:58:30.212146Z","shell.execute_reply.started":"2025-10-05T12:58:30.198733Z","shell.execute_reply":"2025-10-05T12:58:30.211668Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nfrom pathlib import Path\n\n_MR_Z_THRESHOLD = 1.95  # All below this will use 3d method\nseries_path = list(Path(\"/kaggle/input/rsna-niftii-raw/mra_niftii/mra_processed\").iterdir())\n\nvalid_series = []\nfor series in series_path:\n    with open(series / \"metadata.json\", \"r\") as f:\n        metadata = json.load(f)\n\n    spacing = metadata[\"spacing\"][:3]\n    if spacing[-1] <= _MR_Z_THRESHOLD:\n        valid_series.append(str(series))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:30.21267Z","iopub.execute_input":"2025-10-05T12:58:30.212842Z","iopub.status.idle":"2025-10-05T12:58:34.998532Z","shell.execute_reply.started":"2025-10-05T12:58:30.212828Z","shell.execute_reply":"2025-10-05T12:58:34.997953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\ndef main(volume_paths: str, gpu_id: str, output_path: str):\n    import torch\n    import time\n    import gc\n    from inference.inferer import Inferer, generate_and_write_segmentation\n    import cupy as cp\n\n    cp.cuda.Device(gpu_id).use()\n    \n    print(f\"Received {len(volume_paths)} volumes for inference on GPU id: {gpu_id}.\")\n\n        \n    _MR_SPACING_3D = (0.43, 0.43, 0.60)  # MRA 3D\n    _MR_Z_THRESHOLD = 1.95  # All below this will use 3d method\n\n\n    inferer = Inferer(\n        model_save_path=\"/kaggle/input/rsna-niftii-raw/greedy_soup.pth\",\n        inference_dtype=torch.bfloat16,\n        inference_roi_sizes=[(40, 144, 144), (48, 160, 160)],\n        inference_overlap=(0.75, 0.625, 0.625),\n        inference_batch_size=21,\n        use_amp=True,\n        device=f\"cuda:{gpu_id}\",\n    )\n\n    for i, path in enumerate(volume_paths):\n        volume_path = Path(path) / \"volume.nii\"\n        seg_path = Path(output_path) / volume_path.parent.name\n        seg_path.mkdir(parents=True, exist_ok=True)\n        seg_path = seg_path / \"segmentation.nii.gz\"\n        \n        start_time = time.time()\n        generate_and_write_segmentation(\n            volume_path=str(volume_path),\n            inferer=inferer,\n            output_path=str(seg_path),\n            target_spacing=_MR_SPACING_3D,\n        )\n\n        print(f\"Finished {volume_path.parent.name} in {time.time() - start_time} seconds.\")\n        torch.cuda.empty_cache()\n        gc.collect()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:34.999202Z","iopub.execute_input":"2025-10-05T12:58:34.999396Z","iopub.status.idle":"2025-10-05T12:58:35.005302Z","shell.execute_reply.started":"2025-10-05T12:58:34.99938Z","shell.execute_reply":"2025-10-05T12:58:35.004833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import multiprocessing as mp\nimport torch\n\nNUM_GPUS = 4\nTOTAL_SERIES = len(valid_series)\n\ndef split_list_n_parts(lst, n):\n    \"\"\"Split list into n equal parts.\"\"\"\n    k, m = divmod(len(lst), n)\n    return [lst[i*k + min(i, m):(i+1)*k + min(i+1, m)] for i in range(n)]\n\nprint(f\"Splitting {TOTAL_SERIES} series between {NUM_GPUS} gpus\")\n\nchunks = split_list_n_parts(valid_series, NUM_GPUS)\nGPU_IDS = [\"0\", \"1\", \"2\", \"3\"]\n\nprint(\"total chunks: \", len(chunks))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:35.005851Z","iopub.execute_input":"2025-10-05T12:58:35.006003Z","iopub.status.idle":"2025-10-05T12:58:37.94212Z","shell.execute_reply.started":"2025-10-05T12:58:35.00599Z","shell.execute_reply":"2025-10-05T12:58:37.941563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"process_list = []\n\nfor chunk, gpu_id in zip(chunks, GPU_IDS):\n    p = mp.Process(target=main, args=(chunk, gpu_id, \"./mr_segmentation\"))\n    p.start()\n    process_list.append(p)\n\nfor p in process_list:\n    p.join()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-05T12:58:37.943397Z","iopub.execute_input":"2025-10-05T12:58:37.943658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}