From 16979ba25da219eef0190d0df9e5b5f29155f7e4 Mon Sep 17 00:00:00 2001 From: Sathiesh Date: Fri, 4 Sep 2026 11:39:20 +0200 Subject: [PATCH 1/2] Support patch padding and add VS model variants --- fastMONAI/_modidx.py | 20 + fastMONAI/vision_patch.py | 204 ++++++++- nbs/10_vision_patch.ipynb | 395 +++++++++++++++++- research/vestibular_schwannoma/README.md | 6 + .../01_five_fold_cross_validation.ipynb | 4 +- .../tests/workflow/test_models.py | 61 ++- .../tests/workflow/test_train_5fold.py | 7 + .../vestibular_schwannoma/workflow/models.py | 24 +- 8 files changed, 672 insertions(+), 49 deletions(-) diff --git a/fastMONAI/_modidx.py b/fastMONAI/_modidx.py index fee912b..10b7de7 100644 --- a/fastMONAI/_modidx.py +++ b/fastMONAI/_modidx.py @@ -669,6 +669,18 @@ 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch.PatchInferenceEngine.to': ( 'vision_patch.html#patchinferenceengine.to', 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._PadToPatchSize': ( 'vision_patch.html#_padtopatchsize', + 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._PadToPatchSize.__call__': ( 'vision_patch.html#_padtopatchsize.__call__', + 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._PadToPatchSize.__init__': ( 'vision_patch.html#_padtopatchsize.__init__', + 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._PadToPatchSize.transforms': ( 'vision_patch.html#_padtopatchsize.transforms', + 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._PatchLabelSampler': ( 'vision_patch.html#_patchlabelsampler', + 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._PatchLabelSampler.get_probability_map': ( 'vision_patch.html#_patchlabelsampler.get_probability_map', + 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch._PreparedSubject': ( 'vision_patch.html#_preparedsubject', 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch._build_pre_patch_tfms': ( 'vision_patch.html#_build_pre_patch_tfms', @@ -681,12 +693,20 @@ 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch._normalize_patch_overlap': ( 'vision_patch.html#_normalize_patch_overlap', 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._pad_subject_to_patch_size': ( 'vision_patch.html#_pad_subject_to_patch_size', + 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch._predict_one': ( 'vision_patch.html#_predict_one', 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch._predict_patch_tta': ( 'vision_patch.html#_predict_patch_tta', 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._project_mask_to_valid_centers': ( 'vision_patch.html#_project_mask_to_valid_centers', + 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._required_patch_padding': ( 'vision_patch.html#_required_patch_padding', + 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch._save_prediction': ( 'vision_patch.html#_save_prediction', 'fastMONAI/vision_patch.py'), + 'fastMONAI.vision_patch._set_dataset_patch_padding': ( 'vision_patch.html#_set_dataset_patch_padding', + 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch._split_df': ('vision_patch.html#_split_df', 'fastMONAI/vision_patch.py'), 'fastMONAI.vision_patch._stash_from_df_metadata': ( 'vision_patch.html#_stash_from_df_metadata', 'fastMONAI/vision_patch.py'), diff --git a/fastMONAI/vision_patch.py b/fastMONAI/vision_patch.py index 56aa2c7..dc97fe0 100644 --- a/fastMONAI/vision_patch.py +++ b/fastMONAI/vision_patch.py @@ -13,6 +13,7 @@ import pandas as pd import numpy as np import warnings +import copy as _copy import matplotlib.pyplot as plt from pathlib import Path from contextlib import nullcontext @@ -67,6 +68,163 @@ def normalize_patch_transforms(tfms: list) -> list: return None return [_extract_tio_transform(t) for t in tfms] + +def _required_patch_padding(spatial_shape, patch_size): + """Return TorchIO padding bounds and slices that recover the original data.""" + spatial_shape = tuple(int(size) for size in spatial_shape) + patch_size = tuple(int(size) for size in patch_size) + if len(spatial_shape) != 3 or len(patch_size) != 3: + raise ValueError("spatial_shape and patch_size must contain three dimensions") + + padding = [] + valid_slices = [] + for size, minimum_size in zip(spatial_shape, patch_size): + missing = max(minimum_size - size, 0) + before = missing // 2 + after = missing - before + padding.extend((before, after)) + valid_slices.append(slice(before, before + size)) + return tuple(padding), tuple(valid_slices) + + +def _pad_subject_to_patch_size(subject, patch_size, padding_mode=0): + """Pad only undersized axes and return exact slices for removing the padding.""" + padding, valid_slices = _required_patch_padding( + subject.spatial_shape, patch_size + ) + if not any(padding): + return subject, padding, valid_slices + + all_images = subject.get_images_dict(intensity_only=False) + intensity_images = subject.get_images_dict(intensity_only=True) + intensity_names = list(intensity_images) + label_names = [name for name in all_images if name not in intensity_images] + + if intensity_names: + subject = tio.Pad( + padding, padding_mode=padding_mode, include=intensity_names, copy=False + )(subject) + if label_names: + subject = tio.Pad( + padding, padding_mode=0, include=label_names, copy=False + )(subject) + return subject, padding, valid_slices + + +class _PadToPatchSize: + """Run existing preprocessing, then pad undersized images before sampling.""" + def __init__(self, patch_size, padding_mode=0, pre_transform=None): + self.patch_size = tuple(int(size) for size in patch_size) + self.padding_mode = padding_mode + self.pre_transform = pre_transform + + @property + def transforms(self): + """Expose ordered stages for inspection without changing execution semantics.""" + if self.pre_transform is None: + return [self] + previous = getattr(self.pre_transform, "transforms", None) + if previous is None: + previous = [self.pre_transform] + return [*list(previous), self] + + def __call__(self, subject): + if self.pre_transform is not None: + subject = self.pre_transform(subject) + subject, _, _ = _pad_subject_to_patch_size( + subject, self.patch_size, self.padding_mode + ) + return subject + + +def _set_dataset_patch_padding(subjects_dataset, patch_size, padding_mode): + """Return a loader-local dataset that pads after its original preprocessing.""" + existing = getattr(subjects_dataset, "_transform", None) + if isinstance(existing, _PadToPatchSize): + existing = existing.pre_transform + transform = _PadToPatchSize( + patch_size, padding_mode, pre_transform=existing + ) + local_dataset = _copy.copy(subjects_dataset) + local_dataset.set_transform(transform) + return local_dataset + + +def _project_mask_to_valid_centers(mask, exact_axes, patch_size): + """Project class presence onto the only legal center of exact-size axes.""" + projected = mask + for axis in sorted(exact_axes, reverse=True): + projected = projected.any(dim=axis + 1) + result = torch.zeros_like(mask, dtype=torch.bool) + center_index = [slice(None)] + center_index.extend( + int(patch_size[axis] // 2) if axis in exact_axes else slice(None) + for axis in range(3) + ) + result[tuple(center_index)] = projected + return result + + +class _PatchLabelSampler(tio.LabelSampler): + """Keep label-biased sampling meaningful when an axis equals patch size.""" + def get_probability_map(self, subject): + label_image = self.get_probability_map_image(subject) + spatial_shape = np.asarray(label_image.spatial_shape) + exact_axes = [ + axis + for axis, (size, patch) in enumerate(zip(spatial_shape, self.patch_size)) + if size == patch + ] + if not exact_axes: + return super().get_probability_map(subject) + + label_map = label_image.data.float() + if self.label_probabilities_dict is None: + foreground = label_map > 0 + if foreground.shape[0] > 1: + foreground = foreground.any(dim=0, keepdim=True) + return _project_mask_to_valid_centers( + foreground, exact_axes, self.patch_size + ) + + probability_map = torch.zeros( + (1, *spatial_shape.tolist()), + dtype=torch.float32, + device=label_map.device, + ) + crop_ini = self.patch_size // 2 + crop_fin = (self.patch_size - 1) // 2 + valid_slices = tuple( + slice(int(start), int(size - end)) + for start, size, end in zip(crop_ini, spatial_shape, crop_fin) + ) + valid_index = (slice(None), *valid_slices) + label_probabilities = torch.tensor( + list(self.label_probabilities_dict.values()), + dtype=torch.float32, device=label_map.device, + ) + label_probabilities /= label_probabilities.sum() + multichannel = label_map.shape[0] > 1 + items = zip( + self.label_probabilities_dict, label_probabilities, strict=True + ) + for label, label_probability in items: + if multichannel: + class_mask = label_map[label:label + 1].bool() + else: + class_mask = label_map == label + class_mask = _project_mask_to_valid_centers( + class_mask, exact_axes, self.patch_size + ) + valid_mask = class_mask[valid_index] + label_size = valid_mask.sum() + if not label_size: + continue + probability_map[valid_index] += ( + label_probability / label_size * valid_mask + ) + return probability_map + # %% ../nbs/10_vision_patch.ipynb #cell-5 _UNET_DIVISOR = 16 # U-Net-style encoders require patch dims divisible by 2^4 @@ -96,8 +254,10 @@ class PatchConfig: pre_patch_tfms (e.g., normalization) since they were already applied. Inference is unaffected and always applies pre_inference_tfms to raw images. Defaults to False. - padding_mode: Padding mode for CropOrPad when image < patch_size. Default is 0 (zero padding). - Can be int, float, or string (e.g., 'minimum', 'mean'). + padding_mode: Padding mode for intensity images on axes smaller than patch_size, + during both training and inference. Masks are always padded with background 0. + Padding is removed from inference outputs before native-space restoration. + Can be int, float, or string (e.g., 'minimum', 'mean'). Defaults to 0. keep_largest_component: If True, keep only the largest connected component in binary segmentation predictions. Only applies during inference when return_probabilities=False. Defaults to False. Binary-only: for multi-class @@ -345,7 +505,7 @@ def create_patch_sampler(config: PatchConfig) -> tio.data.PatchSampler: return tio.UniformSampler(patch_size) elif config.sampler_type == 'label': - return tio.LabelSampler( + return _PatchLabelSampler( patch_size, label_name='mask', label_probabilities=config.label_probabilities @@ -395,6 +555,9 @@ def __init__( if batch_size <= 0: raise ValueError(f"batch_size must be positive, got {batch_size}") + subjects_dataset = _set_dataset_patch_padding( + subjects_dataset, config.patch_size, config.padding_mode + ) self.subjects_dataset = subjects_dataset self.config = config self.bs = batch_size @@ -1050,6 +1213,8 @@ class _PreparedSubject: org_img: tio.Image input_img: tio.Image org_size: tuple + padding: tuple + valid_slices: tuple grid_sampler: tio.GridSampler aggregator: tio.GridAggregator patch_loader: DataLoader @@ -1174,20 +1339,24 @@ def _prepare_subject(self, img_path: Path | str) -> _PreparedSubject: if self.pre_inference_tfms is not None: subject = self.pre_inference_tfms(subject) - # Pad dimensions smaller than patch_size, keep larger dimensions intact - img_shape = subject['image'].shape[1:] # Exclude channel dim - target_size = [max(s, p) for s, p in zip(img_shape, self.config.patch_size)] + # Pad only undersized dimensions and retain exact slices for restoration. + img_shape = tuple(subject['image'].spatial_shape) + subject, padding, valid_slices = _pad_subject_to_patch_size( + subject, self.config.patch_size, self.config.padding_mode + ) - if any(s < p for s, p in zip(img_shape, self.config.patch_size)): - padded_dims = [f"dim{i}: {s}<{p}" for i, (s, p) in enumerate(zip(img_shape, self.config.patch_size)) if s < p] + if any(padding): + padded_dims = [ + f"dim{i}: {size}<{patch}" + for i, (size, patch) in enumerate(zip(img_shape, self.config.patch_size)) + if size < patch + ] warnings.warn( f"Image size {list(img_shape)} smaller than patch_size {self.config.patch_size} " f"in {padded_dims}. Padding with mode={self.config.padding_mode}. " "Ensure training data covered similar sizes to avoid artifacts." ) - subject = tio.CropOrPad(target_size, padding_mode=self.config.padding_mode)(subject) - patch_overlap = _normalize_patch_overlap(self.config.patch_overlap, self.config.patch_size) grid_sampler = tio.GridSampler( @@ -1200,8 +1369,9 @@ def _prepare_subject(self, img_path: Path | str) -> _PreparedSubject: return _PreparedSubject( subject=subject, org_img=org_img, input_img=input_img, - org_size=org_size, grid_sampler=grid_sampler, - aggregator=aggregator, patch_loader=patch_loader + org_size=org_size, padding=padding, valid_slices=valid_slices, + grid_sampler=grid_sampler, aggregator=aggregator, + patch_loader=patch_loader ) def _run_inference(self, prepared: _PreparedSubject, tta: bool = False) -> torch.Tensor: @@ -1243,6 +1413,16 @@ def _postprocess( return_probabilities keeps the probability map and restores it with linear interpolation. Decoded labels are restored with nearest-neighbor interpolation. """ + valid_slices = getattr(prepared, "valid_slices", None) + if valid_slices is not None: + output = output[(slice(None), *valid_slices)] + expected_shape = tuple(prepared.input_img.spatial_shape) + if tuple(output.shape[1:]) != expected_shape: + raise RuntimeError( + f"Restored inference shape {tuple(output.shape[1:])} does not match " + f"the pre-padding shape {expected_shape}" + ) + if return_probabilities: result = output else: diff --git a/nbs/10_vision_patch.ipynb b/nbs/10_vision_patch.ipynb index bd6dd6f..48c3dab 100644 --- a/nbs/10_vision_patch.ipynb +++ b/nbs/10_vision_patch.ipynb @@ -34,6 +34,7 @@ "import pandas as pd\n", "import numpy as np\n", "import warnings\n", + "import copy as _copy\n", "import matplotlib.pyplot as plt\n", "from pathlib import Path\n", "from contextlib import nullcontext\n", @@ -94,7 +95,164 @@ " \"\"\"\n", " if tfms is None:\n", " return None\n", - " return [_extract_tio_transform(t) for t in tfms]" + " return [_extract_tio_transform(t) for t in tfms]\n", + "\n", + "\n", + "def _required_patch_padding(spatial_shape, patch_size):\n", + " \"\"\"Return TorchIO padding bounds and slices that recover the original data.\"\"\"\n", + " spatial_shape = tuple(int(size) for size in spatial_shape)\n", + " patch_size = tuple(int(size) for size in patch_size)\n", + " if len(spatial_shape) != 3 or len(patch_size) != 3:\n", + " raise ValueError(\"spatial_shape and patch_size must contain three dimensions\")\n", + "\n", + " padding = []\n", + " valid_slices = []\n", + " for size, minimum_size in zip(spatial_shape, patch_size):\n", + " missing = max(minimum_size - size, 0)\n", + " before = missing // 2\n", + " after = missing - before\n", + " padding.extend((before, after))\n", + " valid_slices.append(slice(before, before + size))\n", + " return tuple(padding), tuple(valid_slices)\n", + "\n", + "\n", + "def _pad_subject_to_patch_size(subject, patch_size, padding_mode=0):\n", + " \"\"\"Pad only undersized axes and return exact slices for removing the padding.\"\"\"\n", + " padding, valid_slices = _required_patch_padding(\n", + " subject.spatial_shape, patch_size\n", + " )\n", + " if not any(padding):\n", + " return subject, padding, valid_slices\n", + "\n", + " all_images = subject.get_images_dict(intensity_only=False)\n", + " intensity_images = subject.get_images_dict(intensity_only=True)\n", + " intensity_names = list(intensity_images)\n", + " label_names = [name for name in all_images if name not in intensity_images]\n", + "\n", + " if intensity_names:\n", + " subject = tio.Pad(\n", + " padding, padding_mode=padding_mode, include=intensity_names, copy=False\n", + " )(subject)\n", + " if label_names:\n", + " subject = tio.Pad(\n", + " padding, padding_mode=0, include=label_names, copy=False\n", + " )(subject)\n", + " return subject, padding, valid_slices\n", + "\n", + "\n", + "class _PadToPatchSize:\n", + " \"\"\"Run existing preprocessing, then pad undersized images before sampling.\"\"\"\n", + " def __init__(self, patch_size, padding_mode=0, pre_transform=None):\n", + " self.patch_size = tuple(int(size) for size in patch_size)\n", + " self.padding_mode = padding_mode\n", + " self.pre_transform = pre_transform\n", + "\n", + " @property\n", + " def transforms(self):\n", + " \"\"\"Expose ordered stages for inspection without changing execution semantics.\"\"\"\n", + " if self.pre_transform is None:\n", + " return [self]\n", + " previous = getattr(self.pre_transform, \"transforms\", None)\n", + " if previous is None:\n", + " previous = [self.pre_transform]\n", + " return [*list(previous), self]\n", + "\n", + " def __call__(self, subject):\n", + " if self.pre_transform is not None:\n", + " subject = self.pre_transform(subject)\n", + " subject, _, _ = _pad_subject_to_patch_size(\n", + " subject, self.patch_size, self.padding_mode\n", + " )\n", + " return subject\n", + "\n", + "\n", + "def _set_dataset_patch_padding(subjects_dataset, patch_size, padding_mode):\n", + " \"\"\"Return a loader-local dataset that pads after its original preprocessing.\"\"\"\n", + " existing = getattr(subjects_dataset, \"_transform\", None)\n", + " if isinstance(existing, _PadToPatchSize):\n", + " existing = existing.pre_transform\n", + " transform = _PadToPatchSize(\n", + " patch_size, padding_mode, pre_transform=existing\n", + " )\n", + " local_dataset = _copy.copy(subjects_dataset)\n", + " local_dataset.set_transform(transform)\n", + " return local_dataset\n", + "\n", + "\n", + "def _project_mask_to_valid_centers(mask, exact_axes, patch_size):\n", + " \"\"\"Project class presence onto the only legal center of exact-size axes.\"\"\"\n", + " projected = mask\n", + " for axis in sorted(exact_axes, reverse=True):\n", + " projected = projected.any(dim=axis + 1)\n", + " result = torch.zeros_like(mask, dtype=torch.bool)\n", + " center_index = [slice(None)]\n", + " center_index.extend(\n", + " int(patch_size[axis] // 2) if axis in exact_axes else slice(None)\n", + " for axis in range(3)\n", + " )\n", + " result[tuple(center_index)] = projected\n", + " return result\n", + "\n", + "\n", + "class _PatchLabelSampler(tio.LabelSampler):\n", + " \"\"\"Keep label-biased sampling meaningful when an axis equals patch size.\"\"\"\n", + " def get_probability_map(self, subject):\n", + " label_image = self.get_probability_map_image(subject)\n", + " spatial_shape = np.asarray(label_image.spatial_shape)\n", + " exact_axes = [\n", + " axis\n", + " for axis, (size, patch) in enumerate(zip(spatial_shape, self.patch_size))\n", + " if size == patch\n", + " ]\n", + " if not exact_axes:\n", + " return super().get_probability_map(subject)\n", + "\n", + " label_map = label_image.data.float()\n", + " if self.label_probabilities_dict is None:\n", + " foreground = label_map > 0\n", + " if foreground.shape[0] > 1:\n", + " foreground = foreground.any(dim=0, keepdim=True)\n", + " return _project_mask_to_valid_centers(\n", + " foreground, exact_axes, self.patch_size\n", + " )\n", + "\n", + " probability_map = torch.zeros(\n", + " (1, *spatial_shape.tolist()),\n", + " dtype=torch.float32,\n", + " device=label_map.device,\n", + " )\n", + " crop_ini = self.patch_size // 2\n", + " crop_fin = (self.patch_size - 1) // 2\n", + " valid_slices = tuple(\n", + " slice(int(start), int(size - end))\n", + " for start, size, end in zip(crop_ini, spatial_shape, crop_fin)\n", + " )\n", + " valid_index = (slice(None), *valid_slices)\n", + " label_probabilities = torch.tensor(\n", + " list(self.label_probabilities_dict.values()),\n", + " dtype=torch.float32, device=label_map.device,\n", + " )\n", + " label_probabilities /= label_probabilities.sum()\n", + " multichannel = label_map.shape[0] > 1\n", + " items = zip(\n", + " self.label_probabilities_dict, label_probabilities, strict=True\n", + " )\n", + " for label, label_probability in items:\n", + " if multichannel:\n", + " class_mask = label_map[label:label + 1].bool()\n", + " else:\n", + " class_mask = label_map == label\n", + " class_mask = _project_mask_to_valid_centers(\n", + " class_mask, exact_axes, self.patch_size\n", + " )\n", + " valid_mask = class_mask[valid_index]\n", + " label_size = valid_mask.sum()\n", + " if not label_size:\n", + " continue\n", + " probability_map[valid_index] += (\n", + " label_probability / label_size * valid_mask\n", + " )\n", + " return probability_map" ] }, { @@ -194,8 +352,10 @@ " pre_patch_tfms (e.g., normalization) since they were already applied.\n", " Inference is unaffected and always applies pre_inference_tfms to raw\n", " images. Defaults to False.\n", - " padding_mode: Padding mode for CropOrPad when image < patch_size. Default is 0 (zero padding).\n", - " Can be int, float, or string (e.g., 'minimum', 'mean').\n", + " padding_mode: Padding mode for intensity images on axes smaller than patch_size,\n", + " during both training and inference. Masks are always padded with background 0.\n", + " Padding is removed from inference outputs before native-space restoration.\n", + " Can be int, float, or string (e.g., 'minimum', 'mean'). Defaults to 0.\n", " keep_largest_component: If True, keep only the largest connected component\n", " in binary segmentation predictions. Only applies during inference when\n", " return_probabilities=False. Defaults to False. Binary-only: for multi-class\n", @@ -533,7 +693,7 @@ " return tio.UniformSampler(patch_size)\n", " \n", " elif config.sampler_type == 'label':\n", - " return tio.LabelSampler(\n", + " return _PatchLabelSampler(\n", " patch_size,\n", " label_name='mask',\n", " label_probabilities=config.label_probabilities\n", @@ -643,6 +803,9 @@ " if batch_size <= 0:\n", " raise ValueError(f\"batch_size must be positive, got {batch_size}\")\n", "\n", + " subjects_dataset = _set_dataset_patch_padding(\n", + " subjects_dataset, config.patch_size, config.padding_mode\n", + " )\n", " self.subjects_dataset = subjects_dataset\n", " self.config = config\n", " self.bs = batch_size\n", @@ -1367,6 +1530,8 @@ " org_img: tio.Image\n", " input_img: tio.Image\n", " org_size: tuple\n", + " padding: tuple\n", + " valid_slices: tuple\n", " grid_sampler: tio.GridSampler\n", " aggregator: tio.GridAggregator\n", " patch_loader: DataLoader\n", @@ -1491,20 +1656,24 @@ " if self.pre_inference_tfms is not None:\n", " subject = self.pre_inference_tfms(subject)\n", "\n", - " # Pad dimensions smaller than patch_size, keep larger dimensions intact\n", - " img_shape = subject['image'].shape[1:] # Exclude channel dim\n", - " target_size = [max(s, p) for s, p in zip(img_shape, self.config.patch_size)]\n", + " # Pad only undersized dimensions and retain exact slices for restoration.\n", + " img_shape = tuple(subject['image'].spatial_shape)\n", + " subject, padding, valid_slices = _pad_subject_to_patch_size(\n", + " subject, self.config.patch_size, self.config.padding_mode\n", + " )\n", "\n", - " if any(s < p for s, p in zip(img_shape, self.config.patch_size)):\n", - " padded_dims = [f\"dim{i}: {s}<{p}\" for i, (s, p) in enumerate(zip(img_shape, self.config.patch_size)) if s < p]\n", + " if any(padding):\n", + " padded_dims = [\n", + " f\"dim{i}: {size}<{patch}\"\n", + " for i, (size, patch) in enumerate(zip(img_shape, self.config.patch_size))\n", + " if size < patch\n", + " ]\n", " warnings.warn(\n", " f\"Image size {list(img_shape)} smaller than patch_size {self.config.patch_size} \"\n", " f\"in {padded_dims}. Padding with mode={self.config.padding_mode}. \"\n", " \"Ensure training data covered similar sizes to avoid artifacts.\"\n", " )\n", "\n", - " subject = tio.CropOrPad(target_size, padding_mode=self.config.padding_mode)(subject)\n", - "\n", " patch_overlap = _normalize_patch_overlap(self.config.patch_overlap, self.config.patch_size)\n", "\n", " grid_sampler = tio.GridSampler(\n", @@ -1517,8 +1686,9 @@ "\n", " return _PreparedSubject(\n", " subject=subject, org_img=org_img, input_img=input_img,\n", - " org_size=org_size, grid_sampler=grid_sampler,\n", - " aggregator=aggregator, patch_loader=patch_loader\n", + " org_size=org_size, padding=padding, valid_slices=valid_slices,\n", + " grid_sampler=grid_sampler, aggregator=aggregator,\n", + " patch_loader=patch_loader\n", " )\n", "\n", " def _run_inference(self, prepared: _PreparedSubject, tta: bool = False) -> torch.Tensor:\n", @@ -1560,6 +1730,16 @@ " return_probabilities keeps the probability map and restores it with linear\n", " interpolation. Decoded labels are restored with nearest-neighbor interpolation.\n", " \"\"\"\n", + " valid_slices = getattr(prepared, \"valid_slices\", None)\n", + " if valid_slices is not None:\n", + " output = output[(slice(None), *valid_slices)]\n", + " expected_shape = tuple(prepared.input_img.spatial_shape)\n", + " if tuple(output.shape[1:]) != expected_shape:\n", + " raise RuntimeError(\n", + " f\"Restored inference shape {tuple(output.shape[1:])} does not match \"\n", + " f\"the pre-padding shape {expected_shape}\"\n", + " )\n", + "\n", " if return_probabilities:\n", " result = output\n", " else:\n", @@ -2098,7 +2278,8 @@ " _rows.append({'img': _ipath, 'mask': _mpath, 'is_val': _i >= 2})\n", " _df = pd.DataFrame(_rows)\n", " _cfg = PatchConfig(patch_size=[16, 16, 16], samples_per_volume=2,\n", - " sampler_type='label', label_probabilities={0: 0.5, 1: 0.5})\n", + " sampler_type='label', label_probabilities={0: 0.5, 1: 0.5},\n", + " queue_num_workers=0)\n", "\n", " _dls = MedPatchDataLoaders.from_df(_df, img_col='img', mask_col='mask',\n", " valid_pct=0.5, patch_config=_cfg, seed=0, bs=1)\n", @@ -2120,17 +2301,22 @@ "\n", " # preprocessed=True skips reorder/resample\n", " _cfg_pp = PatchConfig(patch_size=[16, 16, 16], samples_per_volume=2,\n", - " preprocessed=True, target_spacing=[1, 1, 1])\n", + " preprocessed=True, target_spacing=[1, 1, 1],\n", + " queue_num_workers=0)\n", " _dls3 = MedPatchDataLoaders.from_df(_df, img_col='img', mask_col='mask',\n", " valid_pct=0.5, patch_config=_cfg_pp, seed=0, bs=1)\n", " _names_pp = _tfm_names(_dls3.train_ds)\n", " assert 'Resample' not in _names_pp and 'ToCanonical' not in _names_pp, f'preprocessed should skip, got {_names_pp}'\n", + " assert '_PadToPatchSize' in _names_pp, f'preprocessed data still needs minimum-size padding, got {_names_pp}'\n", "\n", " # contrast: non-preprocessed WITH target_spacing includes Resample\n", - " _cfg_rs = PatchConfig(patch_size=[16, 16, 16], samples_per_volume=2, target_spacing=[1, 1, 1])\n", + " _cfg_rs = PatchConfig(patch_size=[16, 16, 16], samples_per_volume=2,\n", + " target_spacing=[1, 1, 1], queue_num_workers=0)\n", " _dls4 = MedPatchDataLoaders.from_df(_df, img_col='img', mask_col='mask',\n", " valid_pct=0.5, patch_config=_cfg_rs, seed=0, bs=1)\n", " assert 'Resample' in _tfm_names(_dls4.train_ds)\n", + " for _test_dls in (_dls, _dls2, _dls3, _dls4):\n", + " _test_dls.close()\n", "\n", "# --- TH: MedPatchDataLoader.__iter__ with patch_tfms yields (MedImage, MedMask) (guards _apply_patch_tfms) ---\n", "with _tempfile.TemporaryDirectory() as _tmp:\n", @@ -2175,6 +2361,183 @@ "print('Safety-net tests (T1-T4, TH, TB) passed!')" ] }, + { + "cell_type": "code", + "execution_count": null, + "id": "minimum-patch-padding-tests", + "metadata": {}, + "outputs": [], + "source": [ + "#| hide\n", + "# Minimum-size padding is exact, mask-safe, and reversible during inference.\n", + "_padding, _valid = _required_patch_padding((16, 16, 49), (16, 16, 64))\n", + "test_eq(_padding, (0, 0, 0, 0, 7, 8))\n", + "test_eq(tuple((axis.start, axis.stop) for axis in _valid), ((0, 16), (0, 16), (7, 56)))\n", + "_even_padding, _ = _required_patch_padding((16, 16, 50), (16, 16, 64))\n", + "test_eq(_even_padding, (0, 0, 0, 0, 7, 7))\n", + "_multi_padding, _ = _required_patch_padding((13, 16, 49), (16, 16, 64))\n", + "test_eq(_multi_padding, (1, 2, 0, 0, 7, 8))\n", + "\n", + "_pad_affine = np.diag([1.0, 1.0, 2.0, 1.0])\n", + "_pad_image = torch.full((1, 16, 16, 49), 5.0)\n", + "_pad_mask = torch.zeros((1, 16, 16, 49), dtype=torch.uint8)\n", + "_pad_mask[0, 8, 8, 24] = 1\n", + "_pad_subject = tio.Subject(\n", + " image=tio.ScalarImage(tensor=_pad_image, affine=_pad_affine),\n", + " mask=tio.LabelMap(tensor=_pad_mask, affine=_pad_affine),\n", + ")\n", + "_padded, _padding, _valid = _pad_subject_to_patch_size(\n", + " _pad_subject, (16, 16, 64), padding_mode='mean'\n", + ")\n", + "test_eq(_padded.spatial_shape, (16, 16, 64))\n", + "assert np.array_equal(_padded.image.affine, _padded.mask.affine)\n", + "assert torch.all(_padded.image.data[..., :7] == 5)\n", + "test_eq(int(_padded.mask.data[..., :7].sum()), 0)\n", + "test_eq(int(_padded.mask.data[..., 56:].sum()), 0)\n", + "test_eq(int(_padded.mask.data.sum()), 1)\n", + "\n", + "_no_pad_subject = tio.Subject(\n", + " image=tio.ScalarImage(tensor=torch.zeros(1, 20, 20, 70), affine=np.eye(4))\n", + ")\n", + "_same_subject, _no_padding, _no_pad_valid = _pad_subject_to_patch_size(\n", + " _no_pad_subject, (16, 16, 64)\n", + ")\n", + "assert _same_subject is _no_pad_subject\n", + "test_eq(_no_padding, (0, 0, 0, 0, 0, 0))\n", + "test_eq(\n", + " tuple((axis.start, axis.stop) for axis in _no_pad_valid),\n", + " ((0, 20), (0, 20), (0, 70)),\n", + ")\n", + "\n", + "# MedPatchDataLoader installs padding after existing preprocessing, before sampling.\n", + "_train_subject = tio.Subject(\n", + " image=tio.ScalarImage(tensor=_pad_image.clone(), affine=_pad_affine),\n", + " mask=tio.LabelMap(tensor=_pad_mask.clone(), affine=_pad_affine),\n", + ")\n", + "_train_dataset = tio.SubjectsDataset([_train_subject])\n", + "_train_config = PatchConfig(\n", + " patch_size=[16, 16, 64],\n", + " samples_per_volume=1,\n", + " sampler_type='label',\n", + " label_probabilities={0: 0.2, 1: 0.8},\n", + " queue_num_workers=0,\n", + ")\n", + "_train_loader = MedPatchDataLoader(_train_dataset, _train_config, batch_size=1)\n", + "_transformed_subject = _train_loader.subjects_dataset[0]\n", + "test_eq(_transformed_subject.spatial_shape, (16, 16, 64))\n", + "_sampled_patch = next(_train_loader.sampler(_transformed_subject, num_patches=1))\n", + "test_eq(_sampled_patch.spatial_shape, (16, 16, 64))\n", + "test_eq(int(_sampled_patch.mask.data.sum()), 1)\n", + "_train_loader.close()\n", + "assert _train_loader.subjects_dataset is not _train_dataset\n", + "test_eq(_train_dataset._transform, None)\n", + "\n", + "# Label sampling projects off-center foreground onto the only valid center plane.\n", + "_off_image = torch.zeros((1, 32, 32, 49))\n", + "_off_mask = torch.zeros((1, 32, 32, 49), dtype=torch.uint8)\n", + "_off_mask[0, 10, 20, 3] = 1\n", + "_off_subject = tio.Subject(\n", + " image=tio.ScalarImage(tensor=_off_image, affine=np.eye(4)),\n", + " mask=tio.LabelMap(tensor=_off_mask, affine=np.eye(4)),\n", + ")\n", + "_off_dataset = tio.SubjectsDataset([_off_subject])\n", + "_off_config = PatchConfig(\n", + " patch_size=[16, 16, 64],\n", + " sampler_type='label',\n", + " label_probabilities={0: 0.2, 1: 0.8},\n", + " samples_per_volume=1,\n", + " queue_num_workers=0,\n", + ")\n", + "_off_loader = MedPatchDataLoader(_off_dataset, _off_config, batch_size=1)\n", + "_off_transformed = _off_loader.subjects_dataset[0]\n", + "_off_probability = _off_loader.sampler.get_probability_map(_off_transformed)\n", + "assert float(_off_probability[0, 10, 20, 32]) > 0\n", + "_off_processed = _off_loader.sampler.process_probability_map(\n", + " _off_probability, _off_transformed\n", + ")\n", + "assert _off_processed[10, 20, 32] > 0\n", + "_default_sampler = _PatchLabelSampler(\n", + " [16, 16, 64], label_name='mask'\n", + ")\n", + "_default_processed = _default_sampler.process_probability_map(\n", + " _default_sampler.get_probability_map(_off_transformed), _off_transformed\n", + ")\n", + "assert _default_processed[10, 20, 32] > 0\n", + "_off_loader.close()\n", + "\n", + "# Existing Compose controls and loader-specific padding remain isolated.\n", + "_controlled = tio.Compose([tio.Clamp(out_min=1, out_max=1)], p=0)\n", + "_shared_dataset = tio.SubjectsDataset([_off_subject], transform=_controlled)\n", + "_cfg64 = PatchConfig(patch_size=[16, 16, 64], queue_num_workers=0)\n", + "_cfg48 = PatchConfig(patch_size=[16, 16, 48], queue_num_workers=0)\n", + "_loader64 = MedPatchDataLoader(_shared_dataset, _cfg64, batch_size=1)\n", + "_loader48 = MedPatchDataLoader(_shared_dataset, _cfg48, batch_size=1)\n", + "assert _shared_dataset._transform is _controlled\n", + "assert _loader64.subjects_dataset is not _loader48.subjects_dataset\n", + "_subject64 = _loader64.subjects_dataset[0]\n", + "_subject48 = _loader48.subjects_dataset[0]\n", + "test_eq(_subject64.spatial_shape, (32, 32, 64))\n", + "test_eq(_subject48.spatial_shape, (32, 32, 49))\n", + "test_eq(float(_subject64.image.data[..., 7:56].sum()), 0.0)\n", + "test_eq(float(_subject48.image.data.sum()), 0.0)\n", + "_loader64.close()\n", + "_loader48.close()\n", + "\n", + "# A no-padding dataset access adds no transform-level RNG draw.\n", + "_no_rng_dataset = tio.SubjectsDataset([_no_pad_subject])\n", + "_no_rng_config = PatchConfig(patch_size=[16, 16, 64], queue_num_workers=0)\n", + "_no_rng_loader = MedPatchDataLoader(_no_rng_dataset, _no_rng_config, batch_size=1)\n", + "_rng_state = torch.random.get_rng_state()\n", + "_ = _no_rng_loader.subjects_dataset[0]\n", + "assert torch.equal(_rng_state, torch.random.get_rng_state())\n", + "_no_rng_loader.close()\n", + "\n", + "# Inference removes the recorded odd padding before decoding or native restoration.\n", + "import tempfile as _padding_tempfile\n", + "import os as _padding_os\n", + "import nibabel as _padding_nib\n", + "\n", + "with _padding_tempfile.TemporaryDirectory() as _tmp:\n", + " _image_path = _padding_os.path.join(_tmp, 'image_49.nii.gz')\n", + " _native_affine = np.diag([1.0, 1.0, 2.0, 1.0])\n", + " _padding_nib.save(\n", + " _padding_nib.Nifti1Image(\n", + " np.zeros((16, 16, 49), dtype=np.float32), _native_affine\n", + " ),\n", + " _image_path,\n", + " )\n", + " _padding_engine = PatchInferenceEngine(\n", + " torch.nn.Conv3d(1, 2, 1),\n", + " PatchConfig(patch_size=[16, 16, 64], apply_reorder=False),\n", + " amp=False,\n", + " )\n", + " _prepared = _padding_engine._prepare_subject(_image_path)\n", + " test_eq(_prepared.padding, (0, 0, 0, 0, 7, 8))\n", + " test_eq(_prepared.subject.spatial_shape, (16, 16, 64))\n", + " _output = torch.zeros(2, 16, 16, 64)\n", + " _output[0] = 1\n", + " _output[0, 8, 8, 27:31] = 0\n", + " _output[1, 8, 8, 27:31] = 2\n", + "\n", + " _restored_mask, _restored_affine = _padding_engine._postprocess(\n", + " _output, _prepared, return_probabilities=False\n", + " )\n", + " test_eq(tuple(_restored_mask.shape), (1, 16, 16, 49))\n", + " test_eq(torch.where(_restored_mask[0, 8, 8] > 0)[0].tolist(), [20, 21, 22, 23])\n", + " assert np.array_equal(_restored_affine, _native_affine)\n", + "\n", + " _restored_probabilities, _ = _padding_engine._postprocess(\n", + " _output, _prepared, return_probabilities=True\n", + " )\n", + " test_eq(tuple(_restored_probabilities.shape), (2, 16, 16, 49))\n", + " test_eq(\n", + " torch.where(_restored_probabilities[1, 8, 8] > 0)[0].tolist(),\n", + " [20, 21, 22, 23],\n", + " )\n", + "\n", + "print('Minimum-size patch padding tests passed!')" + ] + }, { "cell_type": "code", "execution_count": null, diff --git a/research/vestibular_schwannoma/README.md b/research/vestibular_schwannoma/README.md index c21320b..f4c8655 100644 --- a/research/vestibular_schwannoma/README.md +++ b/research/vestibular_schwannoma/README.md @@ -32,12 +32,18 @@ Use the CLI for unattended training: ```bash python train_5fold.py --models unet --folds 1 --epochs 5 --no-compile python train_5fold.py --models unet # One model, all five folds +python train_5fold.py --models dynunet_small dynunet_xs python train_5fold.py --skip-unavailable ``` The default is three models, five folds, and 500 epochs. Models and folds run sequentially within a launcher. Run `python train_5fold.py --help` for all options. +DynUNet variants use the same architecture, deep supervision, patch size, and training +settings. They differ only in channel widths: the existing `dynunet` uses +`[32, 64, 128, 256, 320]`, `dynunet_small` uses `[16, 32, 64, 128, 160]`, and +`dynunet_xs` uses `[8, 16, 32, 64, 80]`. + For one process per GPU, assign one visible GPU and a new `--results-root` to each process: ```bash diff --git a/research/vestibular_schwannoma/notebooks/01_five_fold_cross_validation.ipynb b/research/vestibular_schwannoma/notebooks/01_five_fold_cross_validation.ipynb index f8adb28..36ac684 100644 --- a/research/vestibular_schwannoma/notebooks/01_five_fold_cross_validation.ipynb +++ b/research/vestibular_schwannoma/notebooks/01_five_fold_cross_validation.ipynb @@ -11,11 +11,11 @@ "\n", "## Task and data\n", "\n", - "The dataset used in this study contains 346 contrast-enhanced T1-weighted (ceT1) scans with matching tumour masks. It combines two separate cohorts: 241 cases from the Vestibular-Schwannoma-SEG dataset at Queen Square Radiosurgery Centre in London [[1](https://doi.org/10.1038/s41597-021-01064-w), [2](https://doi.org/10.7937/TCIA.9YTJ-5Q73)], and 105 cases from Elisabeth-TweeSteden Hospital (ETZ) in Tilburg, released for crossMoDA 2022 [[3](https://doi.org/10.5281/zenodo.6504722), [4](https://doi.org/10.1016/j.media.2022.102628)]. CrossMoDA 2022 also contains London cases, but we exclude them to avoid adding the same data twice.\n", + "The dataset used in this study contains 344 contrast-enhanced T1-weighted (ceT1) scans with matching tumour masks. It combines two separate cohorts: 239 cases from the Vestibular-Schwannoma-SEG dataset at Queen Square Radiosurgery Centre in London [[1](https://doi.org/10.1038/s41597-021-01064-w), [2](https://doi.org/10.7937/TCIA.9YTJ-5Q73)], and 105 cases from Elisabeth-TweeSteden Hospital (ETZ) in Tilburg, released for crossMoDA 2022 [[3](https://doi.org/10.5281/zenodo.6504722), [4](https://doi.org/10.1016/j.media.2022.102628)]. CrossMoDA 2022 also contains London cases, but we exclude them to avoid adding the same data twice.\n", "\n", "## Cross-validation and models\n", "\n", - "The fixed `fold` column assigns every case to exactly one of five validation folds (69–70 cases per fold). All models use the same preprocessing, training, and evaluation pipeline." + "The fixed `fold` column assigns every case to exactly one of five validation folds (68–70 cases per fold). All models use the same preprocessing, training, and evaluation pipeline." ] }, { diff --git a/research/vestibular_schwannoma/tests/workflow/test_models.py b/research/vestibular_schwannoma/tests/workflow/test_models.py index 3c6d04d..467aeaa 100644 --- a/research/vestibular_schwannoma/tests/workflow/test_models.py +++ b/research/vestibular_schwannoma/tests/workflow/test_models.py @@ -7,11 +7,16 @@ class TrainingModelConfigTests(unittest.TestCase): def test_specs_are_the_single_architecture_declaration(self): self.assertEqual(models.UNET_SPEC["arch_id"], "monai.unet") - self.assertEqual(models.DYNUNET_SPEC["arch_id"], "monai.dynunet") - self.assertEqual( - models.DYNUNET_SPEC["wrapper_spec"][0]["wrapper_id"], - "fastmonai.dynunet_ds_adapter", - ) + for spec in ( + models.DYNUNET_SPEC, + models.DYNUNET_SMALL_SPEC, + models.DYNUNET_XS_SPEC, + ): + self.assertEqual(spec["arch_id"], "monai.dynunet") + self.assertEqual( + spec["wrapper_spec"][0]["wrapper_id"], + "fastmonai.dynunet_ds_adapter", + ) self.assertEqual( models.UNET_SPEC["arch_kwargs"]["channels"], [32, 64, 128, 256, 320], @@ -20,12 +25,28 @@ def test_specs_are_the_single_architecture_declaration(self): models.DYNUNET_SPEC["arch_kwargs"]["filters"], [32, 64, 128, 256, 320], ) + self.assertEqual( + models.DYNUNET_SMALL_SPEC["arch_kwargs"]["filters"], + [16, 32, 64, 128, 160], + ) + self.assertEqual( + models.DYNUNET_XS_SPEC["arch_kwargs"]["filters"], + [8, 16, 32, 64, 80], + ) self.assertEqual(models.SEGMAMBA_SPEC["arch_id"], "segmamba.v2") + self.assertEqual( + models.SEGMAMBA_SPEC["arch_kwargs"]["feat_size"], + [32, 64, 128, 256], + ) + self.assertEqual( + models.SEGMAMBA_SPEC["arch_kwargs"]["hidden_size"], + 320, + ) def test_registry_contains_only_supported_architectures(self): self.assertEqual( set(models.TRAINING_MODEL_CONFIGS), - {"unet", "dynunet", "segmamba"}, + {"unet", "dynunet", "dynunet_small", "dynunet_xs", "segmamba"}, ) def test_loss_specs_declare_scientifically_relevant_parameters(self): @@ -41,18 +62,26 @@ def test_loss_specs_declare_scientifically_relevant_parameters(self): }, }, ) - self.assertEqual( - models.TRAINING_MODEL_CONFIGS["dynunet"].loss_spec, - { - "loss_id": "monai.deep_supervision", - "kwargs": {"weight_mode": "exp"}, - "base_loss": models.DICE_CE_LOSS_SPEC, - }, - ) + expected = { + "loss_id": "monai.deep_supervision", + "kwargs": {"weight_mode": "exp"}, + "base_loss": models.DICE_CE_LOSS_SPEC, + } + for key in ("dynunet", "dynunet_small", "dynunet_xs"): + with self.subTest(key=key): + self.assertEqual( + models.TRAINING_MODEL_CONFIGS[key].loss_spec, + expected, + ) def test_declared_order_is_preserved(self): - configs = models.get_training_model_configs(("dynunet", "unet")) - self.assertEqual(list(configs), ["dynunet", "unet"]) + configs = models.get_training_model_configs( + ("dynunet_xs", "dynunet_small", "dynunet", "unet") + ) + self.assertEqual( + list(configs), + ["dynunet_xs", "dynunet_small", "dynunet", "unet"], + ) def test_registry_keys_and_conventional_names_are_consistent(self): for key, config in models.TRAINING_MODEL_CONFIGS.items(): diff --git a/research/vestibular_schwannoma/tests/workflow/test_train_5fold.py b/research/vestibular_schwannoma/tests/workflow/test_train_5fold.py index 8f3ba80..4b107f8 100644 --- a/research/vestibular_schwannoma/tests/workflow/test_train_5fold.py +++ b/research/vestibular_schwannoma/tests/workflow/test_train_5fold.py @@ -18,6 +18,13 @@ def test_70_30_sampling_remains_an_explicit_comparison(self): self.assertEqual(args.foreground_probability, 0.7) + def test_smaller_dynunet_variants_are_selectable(self): + args = train_5fold._parser().parse_args( + ["--models", "dynunet_small", "dynunet_xs"] + ) + + self.assertEqual(args.models, ["dynunet_small", "dynunet_xs"]) + if __name__ == "__main__": unittest.main() diff --git a/research/vestibular_schwannoma/workflow/models.py b/research/vestibular_schwannoma/workflow/models.py index 2ebbdd7..797c6ce 100644 --- a/research/vestibular_schwannoma/workflow/models.py +++ b/research/vestibular_schwannoma/workflow/models.py @@ -121,6 +121,8 @@ def _make_dynunet_spec(filters: list[int]) -> dict: DYNUNET_SPEC = _make_dynunet_spec([32, 64, 128, 256, 320]) +DYNUNET_SMALL_SPEC = _make_dynunet_spec([16, 32, 64, 128, 160]) +DYNUNET_XS_SPEC = _make_dynunet_spec([8, 16, 32, 64, 80]) SEGMAMBA_SPEC = make_model_spec( "segmamba.v2", @@ -128,8 +130,8 @@ def _make_dynunet_spec(filters: list[int]) -> dict: "in_chans": 1, "out_chans": 2, "depths": [2, 2, 2, 2], - "feat_size": [48, 96, 192, 384], - "hidden_size": 768, + "feat_size": [32, 64, 128, 256], + "hidden_size": 320, "mamba_backend": "mamba_ssm", }, ) @@ -146,12 +148,28 @@ def _make_dynunet_spec(filters: list[int]) -> dict: ), "dynunet": TrainingModelConfig( key="dynunet", - display_name="DynUNet Small (32-320)", + display_name="DynUNet", model_spec=DYNUNET_SPEC, loss_spec=DYNUNET_LOSS_SPEC, make_loss=_make_dynunet_loss, experiment_name="vestibular_schwannoma_dynunet", ), + "dynunet_small": TrainingModelConfig( + key="dynunet_small", + display_name="DynUNet Small (16-160)", + model_spec=DYNUNET_SMALL_SPEC, + loss_spec=DYNUNET_LOSS_SPEC, + make_loss=_make_dynunet_loss, + experiment_name="vestibular_schwannoma_dynunet_small", + ), + "dynunet_xs": TrainingModelConfig( + key="dynunet_xs", + display_name="DynUNet XS (8-80)", + model_spec=DYNUNET_XS_SPEC, + loss_spec=DYNUNET_LOSS_SPEC, + make_loss=_make_dynunet_loss, + experiment_name="vestibular_schwannoma_dynunet_xs", + ), "segmamba": TrainingModelConfig( key="segmamba", display_name="SegMamba V2", From 1313286fba64644d65e72741b143a82d0972f0e7 Mon Sep 17 00:00:00 2001 From: Sathiesh Date: Fri, 4 Sep 2026 11:58:39 +0200 Subject: [PATCH 2/2] Pin Plum for fastai compatibility --- settings.ini | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/settings.ini b/settings.ini index 55ca3c5..947ade4 100644 --- a/settings.ini +++ b/settings.ini @@ -8,7 +8,7 @@ min_python = 3.10 version = 0.10.1 ### OPTIONAL ### -requirements = fastai>=2.8.6 monai>=1.5.2 torchio>=1.2.1 xlrd>=1.2.0 scikit-image>=0.23.2,<2 imagedata>=3.8.14 mlflow>=3.14 ipython huggingface-hub gdown plum-dispatch safetensors>=0.7.0 +requirements = fastai>=2.8.6 monai>=1.5.2 torchio>=1.2.1 xlrd>=1.2.0 scikit-image>=0.23.2,<2 imagedata>=3.8.14 mlflow>=3.14 ipython huggingface-hub gdown plum-dispatch<2.10 safetensors>=0.7.0 dev_requirements = ipywidgets nbdev<3 execnb<0.2 fastcore<2 tabulate quarto ### nbdev ###