# Copyright (c) 2025 Valentin Boussot
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
# SPDX-License-Identifier: Apache-2.0
"""Tensor and image transforms used in KonfAI preprocessing and postprocessing."""
import itertools
import os
import tempfile
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from enum import Enum
from multiprocessing import current_process, get_context
from pathlib import Path
from typing import Any
import numpy as np
import torch
try:
import SimpleITK as sitk
except ImportError:
sitk = None # type: ignore[assignment]
import torch.nn.functional as F
from konfai import cuda_visible_devices
from konfai.utils.config import _escape_key_component, apply_config
from konfai.utils.dataset import Attribute, Dataset, data_to_image, image_to_data
from konfai.utils.errors import TransformError
from konfai.utils.ITK import _require_simpleitk, box_with_mask, crop_with_mask
from konfai.utils.runtime import NeedDevice
from konfai.utils.utils import get_module, split_path_spec
[docs]
class LocalityKind(Enum):
"""How a transform's output at one voxel depends on its input (its patch-locality contract).
A transform DECLARES its contract via :meth:`Transform.patch_locality`; the patch-streaming
dispatcher (``konfai.data.patching``) reads the declaration and reads only the source region a
target patch actually needs, instead of materialising the whole volume.
- ``POINTWISE`` -- output voxel depends only on the same voxel (and its channels): read the
exact patch.
- ``HALO`` -- bounded neighbourhood: read the patch enlarged by ``halo`` per axis, crop after.
- ``ORIENTATION`` -- flip/permute: read the index-remapped source region.
- ``CROP`` -- the source region is the target region TRANSLATED: reading it IS the answer,
so the stage is not re-applied to it. Unlike a reorientation this drops the voxels outside the
box, so it is no bijection and the stored volume's statistics are not its output's.
- ``GLOBAL_STAT`` -- needs whole-volume stats (``stat_keys`` subset of Min/Max/Mean/Std), obtained
once from disk and cached: read the exact patch + the cached stat.
- ``RESCALE`` -- resample: source region via the scale mapping + interpolation halo.
- ``WHOLE_VOLUME``-- genuinely needs the whole volume: the dispatcher falls back to a full load.
"""
POINTWISE = "pointwise"
HALO = "halo"
ORIENTATION = "orientation"
CROP = "crop"
GLOBAL_STAT = "global_stat"
RESCALE = "rescale"
WHOLE_VOLUME = "whole_volume"
@property
def preserves_statistics(self) -> bool:
"""Whether this kind leaves every whole-volume statistic of its input untouched.
Only a reorientation does: a flip or a permute is a bijection on the voxels, so the multiset of
values -- and therefore Min/Max/Mean/Std over it -- is exactly the input's. Every other kind may
map values (``POINTWISE``, ``GLOBAL_STAT``), mix neighbours (``HALO``) or interpolate
(``RESCALE``). This is what decides whether the statistics of the STORED volume are still those
of a later transform's own input (see ``DatasetManager._plan_stream_region``).
"""
return self is LocalityKind.ORIENTATION
[docs]
@dataclass(frozen=True)
class PatchLocality:
"""A transform's declared patch-locality contract (see :class:`LocalityKind`).
``halo`` is the per-spatial-axis neighbourhood radius in array order (Z, Y, X); a length-1
tuple broadcasts to every axis. ``stat_keys`` are the ``Attribute`` keys a ``GLOBAL_STAT``
transform reads before running (a subset of ``Min``/``Max``/``Mean``/``Std``). ``stat_channels``
restricts the statistic to those channels (``Normalize.channels``).
"""
kind: LocalityKind
halo: tuple[int, ...] = ()
stat_keys: frozenset[str] = field(default_factory=frozenset)
stat_channels: list[int] | None = None
# Overrides the kind-level default (see LocalityKind.preserves_statistics): a POINTWISE transform
# that maps no value (TensorCast to a float dtype) may declare True so a later GLOBAL_STAT can
# still seed from the stored volume.
preserves_statistics: bool | None = None
@property
def statistics_preserving(self) -> bool:
if self.preserves_statistics is not None:
return self.preserves_statistics
return self.kind.preserves_statistics
[docs]
class Foreign(Transform):
"""A transform from another framework, as the loader hands it over.
Name the class where a transform goes and its arguments under it::
transforms:
monai.transforms:ScaleIntensity:
minv: 0.0
maxv: 1.0
The class must be callable on one tensor and return the transformed tensor, which is what
torchvision's transforms, TorchIO's and MONAI's array transforms all are. MONAI's dictionary
transforms (``ScaleIntensityd``) take a dictionary of keys instead: a KonfAI group is the key,
so name the array class.
The class must be DETERMINISTIC: a transform runs on each group of a case in turn, so a random
one would draw again for the label and misalign it from the image. Name it under the
augmentations instead, where a draw is made once for the case and every group is handed it.
It reads the whole volume, which is what a class saying nothing about where its output comes
from is owed. The shape is checked rather than assumed: the patch grid is planned on the shape a
transform announces, and this one announces the shape it was given. Geometry is left as it
stands, which a transform of the intensities alone leaves. A class that resamples, crops or
reorients owns both, and a ``Transform`` subclass is what states them.
"""
def __init__(self, transform, classpath: str) -> None:
super().__init__()
self.classpath = classpath
self.transform = transform
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
result = self.transform(tensor)
if not isinstance(result, torch.Tensor):
result = torch.as_tensor(np.asarray(result))
if list(result.shape) != list(tensor.shape):
raise TransformError(
f"'{self.classpath}' returned the shape {list(result.shape)} for an input of {list(tensor.shape)}.",
"Subclass Transform and implement transform_shape() to declare the shape it returns.",
)
return result
[docs]
class Clip(Transform):
"""Clip tensor intensities to a fixed or data-dependent value range."""
def __init__(
self,
min_value: float | str = -1024,
max_value: float | str = 1024,
save_clip_min: bool = False,
save_clip_max: bool = False,
mask: str | None = None,
) -> None:
super().__init__()
if isinstance(min_value, float) and isinstance(max_value, float) and max_value <= min_value:
raise ValueError(
f"[Clip] Invalid clipping range: max_value ({max_value}) must be greater than min_value ({min_value})"
)
self.min_value = min_value
self.max_value = max_value
self.save_clip_min = save_clip_min
self.save_clip_max = save_clip_max
self.mask = mask
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# A mask reads a separate full volume, and a percentile bound needs the whole histogram:
# both force a whole-volume load. A 'min'/'max' bound needs a global disk statistic
# (GLOBAL_STAT); fixed float bounds clip each voxel independently (POINTWISE).
if self.mask is not None:
return PatchLocality(LocalityKind.WHOLE_VOLUME)
stat_keys: set[str] = set()
for bound, key in ((self.min_value, "Min"), (self.max_value, "Max")):
if isinstance(bound, str):
if bound.lower() == key.lower():
stat_keys.add(key)
else:
return PatchLocality(LocalityKind.WHOLE_VOLUME)
if not stat_keys:
return PatchLocality(LocalityKind.POINTWISE)
return PatchLocality(LocalityKind.GLOBAL_STAT, stat_keys=frozenset(stat_keys))
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
mask = None
if self.mask is not None:
for dataset in self.datasets:
if dataset.is_dataset_exist(self.mask, name):
mask, _ = dataset.read_data(self.mask, name)
break
if mask is None and self.mask is not None:
raise ValueError(
f"Requested mask '{self.mask}' is not present in any dataset. "
"Check your dataset group names or configuration."
)
if mask is None:
tensor_masked = tensor
else:
tensor_masked = tensor[mask == 1]
if isinstance(self.min_value, str):
if self.min_value == "min":
min_value = torch.min(tensor_masked)
elif self.min_value.startswith("percentile:"):
try:
percentile = float(self.min_value.split(":")[1])
# ``np.percentile`` cannot coerce a CUDA tensor (finalize slots may hand Clip a
# GPU-resident volume); ``.cpu()`` is a no-op view on a host tensor.
min_value = np.percentile(tensor_masked.detach().cpu(), percentile)
except (IndexError, ValueError) as exc:
raise ValueError(
f"Invalid format for min_value: '{self.min_value}'. Expected 'percentile:<float>'"
) from exc
else:
raise TypeError(
f"Unsupported string for min_value: '{self.min_value}'."
"Must be a float, 'min', or 'percentile:<float>'."
)
else:
min_value = self.min_value
if isinstance(self.max_value, str):
if self.max_value == "max":
max_value = torch.max(tensor_masked)
elif self.max_value.startswith("percentile:"):
try:
percentile = float(self.max_value.split(":")[1])
max_value = np.percentile(tensor_masked.detach().cpu(), percentile)
except (IndexError, ValueError) as exc:
raise ValueError(
f"Invalid format for max_value: '{self.max_value}'. Expected 'percentile:<float>'"
) from exc
else:
raise TypeError(
f"Unsupported string for max_value: '{self.max_value}'."
" Must be a float, 'max', or 'percentile:<float>'."
)
else:
max_value = self.max_value
# Resolved bounds may be a torch 0-d tensor ("min"/"max") or a numpy scalar
# ("percentile:<p>"); coerce to a Python float so the in-place assignments below are valid
# for a torch tensor across numpy/torch versions.
min_value = float(min_value)
max_value = float(max_value)
# Fast path: one fused in-place clamp instead of two float()-copy + where-scatter passes.
# Restricted to float32 (integer tensors reject float bounds; float16/float64 would compare
# at a different precision than the float()-cast scatter in the else branch below) and to
# non-NaN bounds: a NaN bound — from a dynamic min/max/percentile over data containing NaN —
# makes clamp_ propagate NaN to the whole tensor, whereas the fallback scatter no-ops on it
# (NaN comparisons are False). Every other case takes that fallback, unchanged.
if tensor.dtype == torch.float32 and min_value == min_value and max_value == max_value:
tensor.clamp_(min=min_value, max=max_value)
else:
tensor[torch.where(tensor.float() < min_value)] = min_value
tensor[torch.where(tensor.float() > max_value)] = max_value
if self.save_clip_min:
cache_attribute["Min"] = min_value
if self.save_clip_max:
cache_attribute["Max"] = max_value
return tensor
[docs]
class Normalize(TransformInverse):
"""Map intensities to a target min/max interval and optionally invert it."""
def __init__(
self,
lazy: bool = False,
channels: list[int] | None = None,
min_value: float = -1,
max_value: float = 1,
inverse: bool = True,
) -> None:
super().__init__(inverse)
if max_value <= min_value:
raise ValueError(
f"[Normalize] Invalid range: max_value ({max_value}) must be greater than min_value ({min_value})"
)
self.lazy = lazy
self.min_value = min_value
self.max_value = max_value
self.channels = channels
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Rescaling uses the volume-global Min/Max (restricted to self.channels); the dispatcher reads
# those once from disk and seeds them so every patch (and inverse()) sees the same range.
return PatchLocality(LocalityKind.GLOBAL_STAT, stat_keys=frozenset({"Min", "Max"}), stat_channels=self.channels)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if "Min" not in cache_attribute:
if self.channels:
cache_attribute["Min"] = torch.min(tensor[self.channels])
else:
cache_attribute["Min"] = torch.min(tensor)
if "Max" not in cache_attribute:
if self.channels:
cache_attribute["Max"] = torch.max(tensor[self.channels])
else:
cache_attribute["Max"] = torch.max(tensor)
if not self.lazy:
input_min = float(cache_attribute["Min"])
input_max = float(cache_attribute["Max"])
norm = input_max - input_min
if norm == 0:
print(f"[WARNING] Norm is zero for case '{name}': input is constant with value = {self.min_value}.")
if self.channels:
for channel in self.channels:
tensor[channel].fill_(self.min_value)
else:
tensor.fill_(self.min_value)
else:
if self.channels:
for channel in self.channels:
tensor[channel] = (self.max_value - self.min_value) * (
tensor[channel] - input_min
) / norm + self.min_value
else:
tensor = (self.max_value - self.min_value) * (tensor - input_min) / norm + self.min_value
return tensor
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if self.lazy:
return tensor
else:
input_min = float(cache_attribute.pop("Min"))
input_max = float(cache_attribute.pop("Max"))
return (tensor - self.min_value) * (input_max - input_min) / (self.max_value - self.min_value) + input_min
[docs]
class UnNormalize(Transform):
def __init__(self, min_value: int = -1024, max_value: int = 3071) -> None:
super().__init__()
self.min_value = min_value
self.max_value = max_value
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return (tensor + 1) / 2 * (self.max_value - self.min_value) + self.min_value
[docs]
class Standardize(TransformInverse):
"""Standardize tensors using cached or computed mean and standard deviation."""
def __init__(
self,
lazy: bool = False,
mean: list[float] | None = None,
std: list[float] | None = None,
mask: str | None = None,
inverse: bool = True,
) -> None:
super().__init__(inverse)
self.lazy = lazy
self.mean = mean
self.std = std
self.mask = mask
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# A mask reads a separate full volume (whole-volume). Any of mean/std left unset is taken from
# a volume-global disk statistic (GLOBAL_STAT); when both are given, the standardization is a
# per-voxel affine map with constant coefficients (POINTWISE).
if self.mask is not None:
return PatchLocality(LocalityKind.WHOLE_VOLUME)
stat_keys: set[str] = set()
if self.mean is None:
stat_keys.add("Mean")
if self.std is None:
stat_keys.add("Std")
if not stat_keys:
return PatchLocality(LocalityKind.POINTWISE)
return PatchLocality(LocalityKind.GLOBAL_STAT, stat_keys=frozenset(stat_keys))
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
mask = None
if self.mask is not None:
for dataset in self.datasets:
if dataset.is_dataset_exist(self.mask, name):
mask, _ = dataset.read_data(self.mask, name)
break
if mask is None and self.mask is not None:
raise ValueError(
f"Requested mask '{self.mask}' is not present in any dataset."
" Check your dataset group names or configuration."
)
if mask is None:
tensor_masked = tensor
else:
tensor_masked = tensor[mask == 1]
if "Mean" not in cache_attribute:
cache_attribute["Mean"] = (
torch.tensor([torch.mean(tensor_masked.type(torch.float32))])
if self.mean is None
else torch.tensor(self.mean)
)
if "Std" not in cache_attribute:
cache_attribute["Std"] = (
torch.tensor([torch.std(tensor_masked.type(torch.float32))])
if self.std is None
else torch.tensor(self.std)
)
if self.lazy:
return tensor
else:
mean = self._broadcast(cache_attribute.get_tensor("Mean").to(tensor.device), tensor)
std = self._broadcast(cache_attribute.get_tensor("Std").to(tensor.device), tensor)
return (tensor - mean) / std
@staticmethod
def _broadcast(stat: torch.Tensor, tensor: torch.Tensor) -> torch.Tensor:
"""Shape a scalar or per-channel statistic to broadcast over a channel-first tensor."""
if stat.numel() > 1:
return stat.reshape(-1, *([1] * (tensor.dim() - 1)))
return stat
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if self.lazy:
return tensor
else:
# The stats parse back as float64 on the CPU; move them to the volume's device (the finalize
# chain runs where the volume was blended, possibly CUDA) and compute in float32 so a
# whole-volume fp16 output is not promoted to a float64 copy.
mean = self._broadcast(cache_attribute.pop_tensor("Mean").to(tensor.device, torch.float32), tensor)
std = self._broadcast(cache_attribute.pop_tensor("Std").to(tensor.device, torch.float32), tensor)
return tensor * std + mean
[docs]
class TensorCast(TransformInverse):
# Wide enough to hold every dtype a volume is read as (int8/int16/uint8/float32) with no value moved.
_VALUE_PRESERVING_DTYPES = frozenset({torch.float32, torch.float64})
def __init__(self, dtype: str = "float32", inverse: bool = True) -> None:
super().__init__(inverse)
self.dtype: torch.dtype = getattr(torch, dtype)
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# The promise is that the stored volume's Min/Max/Mean/Std are still a later GLOBAL_STAT's
# input statistics, and a cast keeps them only where it keeps every value. The dtype a volume
# is stored as is not on its header, so the target is what has to hold whatever that is:
# float32 holds an int16 or a float32 exactly, and float16 holds neither -- it runs out of
# mantissa at 2048, where a CT reaches 3000. An integer cast truncates.
return PatchLocality(
LocalityKind.POINTWISE, preserves_statistics=self.dtype in TensorCast._VALUE_PRESERVING_DTYPES
)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
cache_attribute["dtype"] = str(tensor.dtype).replace("torch.", "")
return tensor.type(self.dtype)
[docs]
@staticmethod
def safe_dtype_cast(dtype_str: str) -> torch.dtype:
try:
return getattr(torch, dtype_str)
except AttributeError as exc:
raise ValueError(f"Unsupported dtype: {dtype_str}") from exc
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensor.to(TensorCast.safe_dtype_cast(cache_attribute.pop("dtype")))
[docs]
class Padding(TransformInverse):
def __init__(self, padding: list[int] = [0, 0, 0, 0, 0, 0], mode: str = "constant", inverse: bool = True) -> None:
super().__init__(inverse)
self.padding = padding
self.mode = mode
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if "Origin" in cache_attribute and "Spacing" in cache_attribute and "Direction" in cache_attribute:
origin = torch.tensor(cache_attribute.get_np_array("Origin"))
matrix = torch.tensor(cache_attribute.get_np_array("Direction").reshape((len(origin), len(origin))))
origin = torch.matmul(origin, matrix)
for dim in range(len(self.padding) // 2):
origin[dim] -= self.padding[dim * 2] * cache_attribute.get_np_array("Spacing")[dim]
cache_attribute["Origin"] = torch.matmul(origin, torch.inverse(matrix))
result = F.pad(
tensor.unsqueeze(0),
tuple(self.padding),
self.mode.split(":")[0],
float(self.mode.split(":")[1]) if len(self.mode.split(":")) == 2 else 0,
).squeeze(0)
return result
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: dict[str, torch.Tensor]) -> torch.Tensor:
if "Origin" in cache_attribute and "Spacing" in cache_attribute and "Direction" in cache_attribute:
cache_attribute.pop("Origin")
slices = [slice(0, shape) for shape in tensor.shape]
for dim in range(len(self.padding) // 2):
slices[-dim - 1] = slice(self.padding[dim * 2], tensor.shape[-dim - 1] - self.padding[dim * 2 + 1])
result = tensor[tuple(slices)]
return result
[docs]
class Squeeze(TransformInverse):
def __init__(self, dim: int, inverse: bool = True) -> None:
super().__init__(inverse)
self.dim = dim
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensor.squeeze(self.dim)
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: dict[str, Any]) -> torch.Tensor:
return tensor.unsqueeze(self.dim)
[docs]
class Resample(TransformInverse, ABC):
def __init__(self, inverse: bool) -> None:
super().__init__(inverse)
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# The source region is derived from the scale mapping (read from cache_attribute['Spacing']
# by the dispatcher); a small interpolation halo is added by resample_source_region.
return PatchLocality(LocalityKind.RESCALE)
def _resample(self, tensor: torch.Tensor, size: list[int]) -> torch.Tensor:
if tensor.dtype == torch.uint8:
mode = "nearest"
elif len(tensor.shape) < 4:
mode = "bilinear"
else:
mode = "trilinear"
# Interpolate in the tensor's own float dtype on CUDA. The model output is float16 and CUDA has
# had Half kernels for every mode for years — upcasting the whole (channels x volume) tensor to
# float32 doubled the memory of a multi-class output resample for no argmax benefit. On the CPU,
# keep the historical float32 compute: Half CPU kernels are missing from older torch releases and
# 1.5.8 always computed this path in float32. Integer inputs (uint8 labels) still need a float
# grid for interpolation.
if not tensor.is_floating_point() or (
tensor.device.type == "cpu" and tensor.dtype in (torch.float16, torch.bfloat16)
):
work = tensor.type(torch.float32)
else:
work = tensor
# Return on the input's device (interpolate preserves it): a CPU input stays on the CPU, a
# GPU-resident output volume stays on the GPU so the whole finalize runs where the volume is.
return F.interpolate(work.unsqueeze(0), size=tuple(size), mode=mode).squeeze(0).type(tensor.dtype)
@abstractmethod
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
pass
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
cache_attribute.pop_np_array("Size")
size_1 = cache_attribute.pop_np_array("Size")
if "Spacing" in cache_attribute:
cache_attribute.pop_np_array("Spacing")
return self._resample(tensor, [int(size) for size in size_1])
# Every patch derives its source coordinates from the same global scale (n_in / n_out, from the
# truncated integer sizes F.interpolate itself uses), which is what makes the streamed patches
# agree with the whole-volume call and with each other across a seam.
def _stream_mode(self, tensor: torch.Tensor) -> str:
if tensor.dtype == torch.uint8:
return "nearest"
return "bilinear" if len(tensor.shape) < 4 else "trilinear"
[docs]
def resample_source_region(
self,
target_slices: tuple[slice, ...],
source_spatial_shape: list[int],
cache_attribute: Attribute,
halo: int = 1,
) -> tuple[list[slice], list[int], list[float], list[int], list[int]]:
"""Map a TARGET-grid patch to the minimal SOURCE region to read.
Returns ``(source_slices, region_starts, scales, n_in, n_out)`` — all in
array axis order (Z, Y, X). The ``halo`` is a pure safety margin (the
formula's ``+2`` already captures the i1 neighbour); nearest needs none.
"""
n_in = [int(s) for s in source_spatial_shape]
n_out = [int(s) for s in self.transform_shape("", "", list(n_in), cache_attribute)]
scales = [n_in[k] / n_out[k] for k in range(len(n_in))]
source_slices: list[slice] = []
region_starts: list[int] = []
for k, sl in enumerate(target_slices):
o0, o1 = sl.start, sl.stop
smin = int(np.floor(scales[k] * (o0 + 0.5) - 0.5))
smax = int(np.floor(scales[k] * ((o1 - 1) + 0.5) - 0.5))
a = max(0, smin - halo)
b = min(n_in[k], smax + 2 + halo)
source_slices.append(slice(a, b))
region_starts.append(a)
return source_slices, region_starts, scales, n_in, n_out
[docs]
def resample_region(
self,
sub_tensor: torch.Tensor,
target_slices: tuple[slice, ...],
region_starts: list[int],
scales: list[float],
n_in: list[int],
) -> torch.Tensor:
"""Interpolate a source sub-region to the target patch extent.
``sub_tensor`` is ``[C, (z, y, x)]`` covering ``source_slices``;
``region_starts`` are the global source indices of its first voxel per
axis. Uses the same global coordinate formula as the whole-volume path,
indexing the sub-region as ``sub[i - region_start]``.
"""
mode = self._stream_mode(sub_tensor)
dev = sub_tensor.device
ndim = len(target_slices)
if mode == "nearest":
out = sub_tensor
for k in range(ndim):
# Take the axis's index map from F.interpolate itself, so streamed nearest picks the
# same source voxel as the whole-volume call for every size ratio.
src = torch.arange(n_in[k], device=dev, dtype=torch.float32).reshape(1, 1, -1)
n_out_k = round(n_in[k] / scales[k])
index = F.interpolate(src, size=n_out_k, mode="nearest").long().flatten()
index = index[target_slices[k].start : target_slices[k].stop] - region_starts[k]
out = out.index_select(k + 1, index)
return out
if not sub_tensor.is_floating_point() or (
sub_tensor.device.type == "cpu" and sub_tensor.dtype in (torch.float16, torch.bfloat16)
):
work = sub_tensor.type(torch.float32)
else:
work = sub_tensor
taps: list[tuple[tuple[torch.Tensor, torch.Tensor], tuple[torch.Tensor, torch.Tensor]]] = []
for k in range(ndim):
o = torch.arange(target_slices[k].start, target_slices[k].stop, device=dev, dtype=work.dtype)
src = torch.clamp(scales[k] * (o + 0.5) - 0.5, min=0.0)
i0 = torch.floor(src).long()
i1 = torch.clamp(i0 + 1, max=n_in[k] - 1)
lam = src - i0.to(work.dtype)
taps.append(((i0 - region_starts[k], 1 - lam), (i1 - region_starts[k], lam)))
out_shape = [work.shape[0]] + [sl.stop - sl.start for sl in target_slices]
out = torch.zeros(out_shape, device=dev, dtype=work.dtype)
for combo in itertools.product(*taps):
gathered = work
weight = torch.ones([1] * (ndim + 1), device=dev, dtype=work.dtype)
for k, (idx, lam) in enumerate(combo):
gathered = gathered.index_select(k + 1, idx)
shape = [1] * (ndim + 1)
shape[k + 1] = -1
weight = weight * lam.reshape(shape)
out += gathered * weight
return out.type(sub_tensor.dtype)
[docs]
@abstractmethod
def write_stream_cache_attribute(self, cache_attribute: Attribute, source_spatial_shape: list[int]) -> None:
"""Record the same 'Spacing'/'Size' stack a whole-volume ``__call__`` would.
Called once per case on the persistent attribute so ``inverse()`` at
prediction time pops exactly what the non-streamed path pushed. Uses the
FULL source shape, never the halo'd sub-region.
"""
[docs]
class ResampleToResolution(Resample):
def __init__(self, spacing: list[float] = [1.0, 1.0, 1.0], inverse: bool = True) -> None:
super().__init__(inverse)
self.spacing = torch.tensor([0 if s < 0 else s for s in spacing])
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
image_spacing = cache_attribute.get_tensor("Spacing")
spacing = self.spacing
resize_factor = torch.tensor(
[
s / i_s if s > 0 else 1.0
for s, i_s in zip(self.spacing, cache_attribute.get_tensor("Spacing"), strict=False)
]
)
cache_attribute["Spacing"] = torch.tensor(
[float(s) if s > 0 else float(i_s) for s, i_s in zip(spacing, image_spacing, strict=False)]
)
cache_attribute["Size"] = np.asarray([int(x) for x in torch.tensor(tensor.shape[1:])])
size = [int(x) for x in (torch.tensor(tensor.shape[1:]) * 1 / resize_factor.flip(0))]
cache_attribute["Size"] = np.asarray(size)
return self._resample(tensor, size)
[docs]
def write_stream_cache_attribute(self, cache_attribute: Attribute, source_spatial_shape: list[int]) -> None:
image_spacing = cache_attribute.get_tensor("Spacing")
spacing = self.spacing
resize_factor = torch.tensor(
[s / i_s if s > 0 else 1.0 for s, i_s in zip(self.spacing, image_spacing, strict=False)]
)
cache_attribute["Spacing"] = torch.tensor(
[float(s) if s > 0 else float(i_s) for s, i_s in zip(spacing, image_spacing, strict=False)]
)
cache_attribute["Size"] = np.asarray([int(x) for x in source_spatial_shape])
size = [int(x) for x in (torch.tensor([int(s) for s in source_spatial_shape]) * 1 / resize_factor.flip(0))]
cache_attribute["Size"] = np.asarray(size)
[docs]
class ResampleToShape(Resample):
def __init__(self, shape: list[float] = [100, 256, 256], inverse: bool = True) -> None:
super().__init__(inverse)
self.shape = torch.tensor([0 if s < 0 else s for s in shape])
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
shape = self.shape.clone()
image_shape = torch.tensor([int(x) for x in torch.tensor(tensor.shape[1:])])
for i, s in enumerate(self.shape):
if s == 0:
shape[i] = image_shape[i]
if "Spacing" in cache_attribute:
cache_attribute["Spacing"] = torch.flip(
image_shape / shape * torch.flip(cache_attribute.get_tensor("Spacing"), dims=[0]),
dims=[0],
)
cache_attribute["Size"] = image_shape
cache_attribute["Size"] = shape
return self._resample(tensor, shape)
[docs]
def write_stream_cache_attribute(self, cache_attribute: Attribute, source_spatial_shape: list[int]) -> None:
shape = self.shape.clone()
image_shape = torch.tensor([int(s) for s in source_spatial_shape])
for i, s in enumerate(self.shape):
if s == 0:
shape[i] = image_shape[i]
if "Spacing" in cache_attribute:
cache_attribute["Spacing"] = torch.flip(
image_shape / shape * torch.flip(cache_attribute.get_tensor("Spacing"), dims=[0]),
dims=[0],
)
cache_attribute["Size"] = image_shape
cache_attribute["Size"] = shape
[docs]
class Mask(Transform):
"""Set everything outside a mask to a constant.
Whole-volume: ``__call__`` does not know where its tensor sits, so it cannot read the matching
region of the mask (``Clip(mask=)`` and ``Standardize(mask=)`` load whole volumes for the same
reason).
"""
def __init__(self, path: str = "./default.mha", value_outside: int = 0) -> None:
super().__init__()
self.path = path
self.value_outside = value_outside
self._cached_mask: torch.Tensor | None = None
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if self.path.endswith(".mha"):
_require_simpleitk()
if self._cached_mask is None:
self._cached_mask = torch.tensor(sitk.GetArrayFromImage(sitk.ReadImage(self.path))).unsqueeze(0)
mask = self._cached_mask
else:
mask = None
for dataset in self.datasets:
if dataset.is_dataset_exist(self.path, name):
mask, _ = dataset.read_data(self.path, name)
break
if mask is None:
raise NameError(f"Mask : {self.path}/{name} not found")
# Index on the tensor's own device so the mask works whether the volume is on CPU or GPU
# (``torch.as_tensor`` keeps a torch mask as-is and wraps a numpy one, moving it to the device).
tensor[torch.as_tensor(mask, device=tensor.device) == 0] = self.value_outside
return tensor
[docs]
class Dilate(Transform):
def __init__(self, dilate: int = 1) -> None:
super().__init__()
if dilate < 0:
raise ValueError(f"[Dilate] 'dilate' must be >= 0, got {dilate}")
self.dilate = dilate
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# A box dilation of radius ``dilate`` spreads foreground by at most ``dilate`` voxels per axis:
# a bounded HALO. At the true border the separable max-pool padding matches the whole-volume
# result once the halo clamps, so seams are byte-identical. Radius 0 is a spatial identity.
if self.dilate == 0:
return PatchLocality(LocalityKind.POINTWISE)
return PatchLocality(LocalityKind.HALO, halo=(self.dilate,))
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if self.dilate == 0:
return tensor
data = (tensor > 0).to(torch.float32)
spatial_dims = data.dim() - 1
d = self.dilate
k = 2 * d + 1
# A cubic (box) structuring element is separable: dilating by a k**n box equals n successive
# 1-D max-pools, one per spatial axis. This is bit-identical to a single k**n max-pool (max is
# associative and the box is the Minkowski sum of 1-D segments) for ~k**(n-1)x fewer comparisons
# — the k**3 dense pool is the dominant cost of the whole-volume mask load.
if spatial_dims == 2:
data = F.max_pool2d(data, kernel_size=(k, 1), stride=1, padding=(d, 0))
data = F.max_pool2d(data, kernel_size=(1, k), stride=1, padding=(0, d))
elif spatial_dims == 3:
data = F.max_pool3d(data, kernel_size=(k, 1, 1), stride=1, padding=(d, 0, 0))
data = F.max_pool3d(data, kernel_size=(1, k, 1), stride=1, padding=(0, d, 0))
data = F.max_pool3d(data, kernel_size=(1, 1, k), stride=1, padding=(0, 0, d))
else:
raise ValueError(
"[Dilate] Unsupported tensor shape for "
f"'{name}': expected [C,H,W] or [C,D,H,W], got {list(tensor.shape)}"
)
return data.to(tensor.dtype)
[docs]
class Sum(Transform):
def __init__(self, dim: int = 0) -> None:
super().__init__()
self.dim = dim
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Pointwise only when reducing the leading channel/model axis (dim 0); a spatial sum spans
# the whole extent, so it falls back to the whole volume.
if self.dim == 0:
return PatchLocality(LocalityKind.POINTWISE)
return PatchLocality(LocalityKind.WHOLE_VOLUME)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if "number_of_channels_per_model" in cache_attribute:
number_of_channels = cache_attribute.pop_tensor("number_of_channels_per_model")
result = tensor[0]
for i, t in enumerate(tensor[1:]):
t[t != 0] += int(number_of_channels[i]) - 1
result += t
return result
else:
return torch.sum(tensor, dim=self.dim).to(tensor.dtype)
[docs]
class MergeLabels(Transform):
"""Merge the per-model argmax label maps of a ``combine: Concat`` ensemble into one global map.
Each model's ``Argmax`` produces a LOCAL class index (``0`` = background). A model's
non-background labels are shifted past every earlier model's foreground classes -- by the
CUMULATIVE sum of the earlier models' foreground counts (``nb_class - 1``) -- so the models'
disjoint label ranges tile a single global label space.
This is the label-space counterpart of ``InferenceStack`` (which averages *same-class*
probability ensembles): use ``MergeLabels`` when the models segment DIFFERENT structures, e.g.
the 5-task TotalSegmentator ensemble (organs / vertebrae / cardiac / muscles / ribs). Requires
``number_of_channels_per_model`` in the attribute (written by the ``Concat`` reduction).
Models are assumed to segment disjoint structures, but boundaries disagree in practice: a voxel
claimed by several models takes the label of the LAST model in ensemble order (adding the global
ids instead would fabricate a label belonging to neither model).
"""
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Merges the leading model axis per voxel; spatial support is a single voxel.
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if "number_of_channels_per_model" not in cache_attribute:
raise TransformError(
"MergeLabels expects a multi-model 'combine: Concat' output: "
"'number_of_channels_per_model' is missing from the attribute.",
)
number_of_channels = cache_attribute.pop_tensor("number_of_channels_per_model")
result = tensor[0].clone()
offset = int(number_of_channels[0]) - 1
for i, t in enumerate(tensor[1:]):
foreground = t != 0
result[foreground] = (t[foreground] + offset).to(result.dtype)
offset += int(number_of_channels[i + 1]) - 1
return result
[docs]
class Gradient(Transform):
def __init__(self, per_dim: bool = False):
super().__init__()
self.per_dim = per_dim
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# First-difference gradient: each output voxel reads its immediate neighbour, a HALO of radius
# 1. The far-edge ConstantPad reproduces the whole-volume border once the halo clamps there.
return PatchLocality(LocalityKind.HALO, halo=(1,))
@staticmethod
def _image_gradient_2d(image: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
dx = image[:, 1:, :] - image[:, :-1, :]
dy = image[:, :, 1:] - image[:, :, :-1]
return torch.nn.ConstantPad2d((0, 0, 0, 1), 0)(dx), torch.nn.ConstantPad2d((0, 1, 0, 0), 0)(dy)
@staticmethod
def _image_gradient_3d(
image: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dx = image[:, 1:, :, :] - image[:, :-1, :, :]
dy = image[:, :, 1:, :] - image[:, :, :-1, :]
dz = image[:, :, :, 1:] - image[:, :, :, :-1]
return (
torch.nn.ConstantPad3d((0, 0, 0, 0, 0, 1), 0)(dx),
torch.nn.ConstantPad3d((0, 0, 0, 1, 0, 0), 0)(dy),
torch.nn.ConstantPad3d((0, 1, 0, 0, 0, 0), 0)(dz),
)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
result = torch.stack(
(Gradient._image_gradient_3d(tensor) if len(tensor.shape) == 4 else Gradient._image_gradient_2d(tensor)),
dim=1,
).squeeze(0)
if not self.per_dim:
result = torch.sigmoid(result * 3)
result = result.norm(dim=0)
result = torch.unsqueeze(result, 0)
return result
[docs]
class Argmax(Transform):
def __init__(self, dim: int = 0) -> None:
super().__init__()
self.dim = dim
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Pointwise ONLY when reducing the channel axis (dim 0). Over a spatial axis the argmax spans
# the whole extent, so a per-patch argmax would diverge -- fall back to the whole volume.
if self.dim == 0:
return PatchLocality(LocalityKind.POINTWISE)
return PatchLocality(LocalityKind.WHOLE_VOLUME)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return torch.argmax(tensor, dim=self.dim).unsqueeze(self.dim)
[docs]
class Softmax(Transform):
def __init__(self, dim: int = 0) -> None:
super().__init__()
self.dim = dim
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Pointwise ONLY when reducing the channel axis (dim 0). Over a spatial axis softmax normalises
# across the whole extent, so a per-patch softmax would diverge -- fall back to the whole volume.
if self.dim == 0:
return PatchLocality(LocalityKind.POINTWISE)
return PatchLocality(LocalityKind.WHOLE_VOLUME)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return torch.softmax(tensor, dim=self.dim)
[docs]
class FlatLabel(Transform):
def __init__(self, labels: list[int] | None = None) -> None:
super().__init__()
self.labels = labels
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
data = torch.zeros_like(tensor)
if self.labels:
for label in self.labels:
data[torch.where(tensor == label)] = 1
else:
data[torch.where(tensor > 0)] = 1
return data
[docs]
class Save(Transform):
def __init__(self, dataset: str, group: str | None = None) -> None:
super().__init__()
self.dataset = dataset
self.group = group
# WHOLE_VOLUME on purpose: a Save still in the chain must WRITE the preprocessed volume, and the
# streamed path never has one to write. (A Save whose cache exists is a source boundary instead.)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensor
[docs]
class Flatten(Transform):
def __init__(self) -> None:
super().__init__()
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensor.flatten()
[docs]
class Permute(TransformInverse):
def __init__(self, dims: str = "1|0|2", inverse: bool = True) -> None:
super().__init__(inverse)
self.dims = [0] + [int(d) + 1 for d in dims.split("|")]
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
return PatchLocality(LocalityKind.ORIENTATION)
[docs]
def stream_region_source(
self,
target_slices: tuple[slice, ...],
source_spatial_shape: list[int],
cache_attribute: Attribute,
) -> list[slice]:
# Output spatial axis k comes from input axis ``self.dims[k + 1] - 1`` (self.dims is
# channel-inclusive). Placing each target slice at its source axis yields the source region
# whose permutation reproduces the target patch exactly.
source_slices = [slice(0, n) for n in source_spatial_shape]
for k, sl in enumerate(target_slices):
source_slices[self.dims[k + 1] - 1] = slice(sl.start, sl.stop)
return source_slices
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensor.permute(tuple(self.dims))
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensor.permute(tuple(np.argsort(self.dims)))
[docs]
class Flip(TransformInverse):
def __init__(self, dims: str = "1|0|2", inverse: bool = True) -> None:
super().__init__(inverse)
self.dims = [int(d) + 1 for d in str(dims).split("|")]
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
return PatchLocality(LocalityKind.ORIENTATION)
[docs]
def stream_region_source(
self,
target_slices: tuple[slice, ...],
source_spatial_shape: list[int],
cache_attribute: Attribute,
) -> list[slice]:
# A flipped spatial axis reads the mirror region ``[n - stop, n - start)``; applying the flip
# to that sub-region reproduces the target patch. Non-flipped axes read the identity region.
source_slices: list[slice] = []
for k, sl in enumerate(target_slices):
n = source_spatial_shape[k]
if (k + 1) in self.dims:
source_slices.append(slice(n - sl.stop, n - sl.start))
else:
source_slices.append(slice(sl.start, sl.stop))
return source_slices
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensor.flip(tuple(self.dims))
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensor.flip(tuple(self.dims))
[docs]
class Canonical(TransformInverse):
"""Reorient a volume onto the canonical (LPS) direction cosines.
An orthogonal reorientation is a signed permutation of the axes: an exact index remap (values only
change place, so whole-volume statistics survive); only an oblique direction is resampled. A remap
that permutes axes transposes the extents it swaps, so ``transform_shape`` folds the patch grid
onto the reoriented shape.
"""
# An orthonormal direction's entries are exactly 0 or +/-1 when it is axis-aligned, but the
# reorientation is a product with an inverse, so it lands within a few double ulps of them.
_AXIS_ALIGNED_ATOL = 1e-9
def __init__(self, inverse: bool = True) -> None:
super().__init__(inverse)
self.canonical_direction = torch.diag(torch.tensor([-1, -1, 1])).to(torch.double)
def _reorientation(self, cache_attribute: Attribute) -> torch.Tensor:
"""The map taking an output coordinate onto the input it comes from, in (x, y, z).
A voxel sits at ``D @ (spacing * index) + origin``, so the map is ``D^-1 @ C`` (with the
target spacing carried along the permutation, see ``_carried``) -- NOT the rotation
``C @ D^-1``, which only agrees where the two commute.
"""
initial_matrix = cache_attribute.get_tensor("Direction").reshape(3, 3).to(torch.double)
return initial_matrix.inverse() @ self.canonical_direction
@classmethod
def _index_remap(cls, reorientation: torch.Tensor) -> list[tuple[int, bool]] | None:
"""Per output SPATIAL axis, the source axis it reads and whether it reads it mirrored.
``reorientation`` maps an output coordinate onto the input it comes from, so it is an exact
remap exactly when it is a signed permutation: output physical axis ``c`` then reads input
physical axis ``r``, backwards where the sign is negative. Anything else mixes axes. Axes are
returned in array order, where physical axis k is array axis ``n - 1 - k``. The test (every
column of L1 norm 1 with peak 1) admits exactly the signed permutations: unit column sums
alone would also pass an axis-averaging matrix.
"""
n = reorientation.shape[0]
unit = torch.ones(n, dtype=reorientation.dtype)
columns = reorientation.abs()
if not torch.allclose(columns.sum(0), unit, atol=cls._AXIS_ALIGNED_ATOL):
return None
if not torch.allclose(columns.amax(0), unit, atol=cls._AXIS_ALIGNED_ATOL):
return None
remap = []
for c in reversed(range(n)):
r = int(columns[:, c].argmax())
remap.append((n - 1 - r, bool(reorientation[r, c] < 0)))
return remap
def _orthogonal_remap(self, cache_attribute: Attribute) -> list[tuple[int, bool]] | None:
"""The exact index remap this case's reorientation is, or ``None`` where it is not one.
Total: a case whose header carries no usable direction cosines has no remap to make, and an
oblique one has none to make either -- both answer ``None`` rather than raise, and the resample
is what answers for them.
"""
if "Direction" not in cache_attribute or cache_attribute.get_np_array("Direction").size != 9:
return None
return Canonical._index_remap(self._reorientation(cache_attribute))
@staticmethod
def _carried(per_physical_axis: torch.Tensor, remap: list[tuple[int, bool]] | None) -> torch.Tensor:
"""Carry a per-physical-axis quantity along a remap: output axis c takes the axis it reads.
A spacing and a half-extent travel with the axis they belong to -- what a reorientation
preserves is the volume's physical extent, not which axis carries it. An oblique direction is
resampled onto the input's own grid, so without a remap nothing moves.
"""
if remap is None:
return per_physical_axis
# The remap is in array order and these are (x, y, z): read in array order, gather, restore.
return per_physical_axis.flip(0)[[source for source, _ in remap]].flip(0)
@staticmethod
def _half_extent(spatial_shape: list[int], spacing: torch.Tensor) -> torch.Tensor:
"""Half a grid's physical extent along each axis, in (x, y, z). A shape is in array order."""
return torch.tensor(
[(spatial_shape[-axis - 1] - 1) * spacing[axis] / 2 for axis in range(len(spatial_shape))],
dtype=torch.double,
)
@staticmethod
def _affine_matrix(matrix: torch.Tensor, translation: torch.Tensor) -> torch.Tensor:
return torch.cat(
(
torch.cat((matrix, translation.unsqueeze(0).T), dim=1),
torch.tensor([[0, 0, 0, 1]]),
),
dim=0,
)
@staticmethod
def _resample_affine(data: torch.Tensor, matrix: torch.Tensor):
if data.dtype == torch.uint8:
mode = "nearest"
else:
mode = "bilinear"
# Sample in the data's own device and float dtype: the model output is float16 on the GPU, and
# affine_grid/grid_sample support float16 on CPU and CUDA. Building the grid on the data's device
# (instead of a CPU float32 grid) keeps the whole reorientation on-device — no host round-trip and
# no float32 upcast of the (channels x volume) tensor. Integer inputs still need a float grid.
# Accepted trade-off: an fp16 grid quantizes the sampling coordinates (up to ~0.1 voxel at 512^3),
# chosen over the ~2x transient memory of a float32 grid + volume upcast.
work = data if data.is_floating_point() else data.type(torch.float32)
grid = torch.nn.functional.affine_grid(
matrix[:, :-1, ...].to(device=work.device, dtype=work.dtype),
[1, *list(data.shape)],
align_corners=True,
)
return (
torch.nn.functional.grid_sample(
work.unsqueeze(0),
grid,
align_corners=True,
mode=mode,
padding_mode="reflection",
)
.squeeze(0)
.type(data.dtype)
)
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Only the case can say which reorientation this is, so only the header can answer. An orthogonal
# one -- mirroring or permuting -- remaps indices, which is what ORIENTATION streams; an oblique
# one is resampled from the whole volume.
if self._orthogonal_remap(cache_attribute) is None:
return PatchLocality(LocalityKind.WHOLE_VOLUME)
return PatchLocality(LocalityKind.ORIENTATION)
[docs]
def stream_region_source(
self,
target_slices: tuple[slice, ...],
source_spatial_shape: list[int],
cache_attribute: Attribute,
) -> list[slice]:
# Target axis k reads source axis ``source``, so the target slice IS the source's -- taken at the
# far end ``[n - stop, n - start)`` where the remap reads that axis backwards. Flipping the region
# read reproduces the patch: a flip restricted to a contiguous region is that region reversed.
# Both the slices and the remap are in array order, and the remap covers every axis exactly once.
remap = self._orthogonal_remap(cache_attribute)
if remap is None:
raise TransformError(
"Canonical declared a region patch-locality for a direction it cannot remap exactly.",
"Report this: patch_locality() and stream_region_source() disagree about the case.",
)
source_slices = [slice(None)] * len(remap)
for target, (source, mirrored) in zip(target_slices, remap, strict=False):
extent = source_spatial_shape[source]
source_slices[source] = (
slice(extent - target.stop, extent - target.start) if mirrored else slice(target.start, target.stop)
)
return source_slices
[docs]
def write_stream_cache_attribute(self, cache_attribute: Attribute, source_spatial_shape: list[int]) -> None:
initial_matrix = cache_attribute.get_tensor("Direction").reshape(3, 3).to(torch.double)
initial_origin = cache_attribute.get_tensor("Origin")
spacing = cache_attribute.get_tensor("Spacing").to(torch.double)
remap = self._orthogonal_remap(cache_attribute)
half_extent = Canonical._half_extent(source_spatial_shape, spacing)
cache_attribute["Direction"] = self.canonical_direction.flatten()
cache_attribute["Spacing"] = Canonical._carried(spacing, remap)
# The reorientation fixes the volume's centre, so the new origin is that centre stepped back by
# the canonical half-extent -- the TARGET grid's, which a permutation has carried onto other
# axes. The extent is the VOLUME's, never a patch's: it is an argument rather than the handed
# tensor's shape.
center = initial_matrix @ half_extent + initial_origin
cache_attribute["Origin"] = center - self.canonical_direction @ Canonical._carried(half_extent, remap)
def _reorient(self, tensor: torch.Tensor, reorientation: torch.Tensor) -> torch.Tensor:
"""Apply a reorientation: an exact index remap where it is one, a resample where it is not.
An orthogonal reorientation is a bijection on the voxels, so it must reproduce the input's
multiset bit for bit -- which only a permute and a flip do.
"""
remap = Canonical._index_remap(reorientation)
if remap is None:
matrix = Canonical._affine_matrix(reorientation, torch.tensor([0, 0, 0]))
return Canonical._resample_affine(tensor, matrix.unsqueeze(0))
# The remap is spatial and the tensor is channel-first, so the channel axes lead it unpermuted.
offset = tensor.dim() - len(remap)
dims = list(range(offset)) + [offset + source for source, _ in remap]
flips = [offset + axis for axis, (_, mirrored) in enumerate(remap) if mirrored]
# flip materialises the permuted view, so the result never aliases the tensor it was read from.
return tensor.permute(dims).flip(flips)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
# Read the source geometry before recording the canonical one over it: the attribute stacks.
reorientation = self._reorientation(cache_attribute)
self.write_stream_cache_attribute(cache_attribute, list(tensor.shape[1:]))
return self._reorient(tensor, reorientation)
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
# Popping restores the source geometry, which is what the inverse remap is then read from.
cache_attribute.pop("Direction")
cache_attribute.pop("Spacing")
cache_attribute.pop("Origin")
return self._reorient(tensor, self._reorientation(cache_attribute).inverse())
[docs]
class HistogramMatching(Transform):
"""Match a volume's intensity distribution onto a reference group's.
Whole-volume: the LUT is built from the volume's 256-bin histogram, which is not a statistic
``GLOBAL_STAT`` names and cannot be read back out of the sitk filter.
"""
def __init__(self, reference_group: str) -> None:
super().__init__()
self.reference_group = reference_group
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
image = data_to_image(tensor, cache_attribute)
image_ref = None
for dataset in self.datasets:
if dataset.is_dataset_exist(self.reference_group, name):
image_ref = dataset.read_image(self.reference_group, name)
if image_ref is None:
raise NameError(f"Image : {self.reference_group}/{name} not found")
_require_simpleitk()
matcher = sitk.HistogramMatchingImageFilter()
matcher.SetNumberOfHistogramLevels(256)
matcher.SetNumberOfMatchPoints(1)
matcher.SetThresholdAtMeanIntensity(True)
result, _ = image_to_data(matcher.Execute(image, image_ref))
return torch.tensor(result)
[docs]
class SelectLabel(Transform):
def __init__(self, labels: list[str]) -> None:
super().__init__()
self.labels = [label[1:-1].split(",") for label in labels]
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
data = torch.zeros_like(tensor)
for old_label, new_label in self.labels:
data[tensor == int(old_label)] = int(new_label)
return data
[docs]
class OneHot(TransformInverse):
def __init__(self, num_classes: int, inverse: bool = True) -> None:
super().__init__(inverse)
self.num_classes = num_classes
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Expands each voxel's scalar label into a one-hot channel vector (spatially pointwise).
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
result = (
F.one_hot(tensor.type(torch.int64), num_classes=self.num_classes)
.permute(0, len(tensor.shape), *[i + 1 for i in range(len(tensor.shape) - 1)])
.float()
.squeeze(0)
)
return result
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
# Argmax the CLASS axis (the one sized num_classes) and re-insert it, restoring a [.., 1, *spatial]
# label map. The predictor feeds this per-sample output[i] = [num_classes, *spatial] (class axis 0),
# but a batched [B, num_classes, *spatial] (class axis 1) is also handled, so it never argmaxes a
# batch or spatial axis.
class_dim = 0 if tensor.shape[0] == self.num_classes else 1
return torch.argmax(tensor, dim=class_dim).unsqueeze(class_dim)
# Published app used by KonfAIInference when the configuration leaves repo/model unset.
DEFAULT_INFERENCE_REPO_ID = "VBoussot/MRSegmentator-KonfAI"
DEFAULT_INFERENCE_MODEL_NAME = "MRSegmentator"
[docs]
class KonfAIInference(Transform):
supports_dataloader_workers = False
def __init__(
self,
repo_id: str = DEFAULT_INFERENCE_REPO_ID,
model_name: str = DEFAULT_INFERENCE_MODEL_NAME,
checkpoints_name: list[str] = ["fold_0"],
number_of_tta: int = 0,
number_of_mc: int = 0,
per_channel: bool = False,
):
super().__init__()
self.repo_id = repo_id
self.model_name = model_name
self.checkpoints_name = checkpoints_name
self.number_of_tta = number_of_tta
self.number_of_mc = number_of_mc
self.per_channel = per_channel
[docs]
def infer_entry(self, dataset_path: Path, output_path: Path, gpu: list[int]):
try:
from konfai_apps import KonfAIApp
except ImportError as exc: # pragma: no cover - depends on optional install
raise RuntimeError(
"KonfAIInference requires the standalone 'konfai-apps' package. "
"Install it from the repository with 'pip install -e ./konfai-apps'."
) from exc
# Nested KonfAI runs must choose their own rendezvous ports instead of
# inheriting the parent's already-bound distributed settings.
os.environ.pop("KONFAI_MASTER_PORT", None)
os.environ.pop("KONFAI_TENSORBOARD_PORT", None)
konfai_app = KonfAIApp(f"{self.repo_id}:{self.model_name}", False, False)
konfai_app.infer(
[[dataset_path]],
output_path,
0,
self.checkpoints_name,
self.number_of_tta,
mc=0,
uncertainty=False,
gpu=gpu,
)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if current_process().daemon:
raise RuntimeError(
"KonfAIInference cannot run inside daemon DataLoader workers. "
"Use 'Dataset.num_workers: 0' for pipelines that include this transform."
)
_require_simpleitk()
with tempfile.TemporaryDirectory() as tmpdir:
dataset_path = Path(tmpdir) / "Dataset"
if self.per_channel:
for i, channel in enumerate(tensor):
image = data_to_image(channel.unsqueeze(0), cache_attribute)
(dataset_path / f"P{i:03d}").mkdir(parents=True, exist_ok=True)
sitk.WriteImage(image, str(dataset_path / f"P{i:03d}" / "Volume.mha"))
else:
image = data_to_image(tensor, cache_attribute)
(dataset_path / "P000").mkdir(parents=True, exist_ok=True)
sitk.WriteImage(image, str(dataset_path / "P000" / "Volume.mha"))
ctx = get_context("spawn")
p = ctx.Process(
target=self.infer_entry, args=(dataset_path, Path(tmpdir) / "Output", cuda_visible_devices())
)
p.start()
p.join()
if p.exitcode != 0:
raise RuntimeError("Inference process failed")
return self._reassemble_output(Path(tmpdir) / "Output")
@staticmethod
def _reassemble_output(output_dir: Path) -> torch.Tensor:
result = []
for file in sorted(output_dir.rglob("*.mha")):
if file.name != "InferenceStack.mha":
result.append(torch.from_numpy(image_to_data(sitk.ReadImage(str(file)))[0]))
return torch.stack(result, dim=1).squeeze(0)
[docs]
class InferenceStack(Transform):
def __init__(self, dataset: str, name: str, mode: str = "mean"):
super().__init__()
self.dataset = None
if dataset:
filename, _, file_format = split_path_spec(dataset)
self.dataset = Dataset(filename, file_format)
self.name = name
self.mode = mode
# patch_locality stays the WHOLE_VOLUME default: the member reduction is pointwise, but __call__
# also WRITES the whole per-member stack to disk (like Save), which a per-patch pass cannot do. It
# is an ensemble/finalize transform, never in a streaming input chain, so this costs no streaming.
def __call__(self, name: str, tensors: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if tensors.shape[0] == 1:
return tensors.squeeze(0)
if self.mode == "Seg":
_tensors = torch.argmax(torch.softmax(tensors, dim=1), dim=1).to(torch.uint8)
else:
_tensors = tensors.squeeze(1)
dataset = self.dataset if self.dataset else self.datasets[-1]
dataset.write("InferenceStack", name, _tensors.float().cpu().numpy(), cache_attribute)
return (
torch.median(tensors.float(), dim=0).values.to(tensors.dtype)
if self.mode == "median"
else tensors.float().mean(0).to(tensors.dtype)
)
[docs]
class Norm(Transform):
"""Vector magnitude over the trailing component axis.
Reduces a stacked vector field (e.g. a displacement-field ensemble ``[N, (D), H, W, C]``) to
per-sample magnitudes ``[N, (D), H, W]``, typically before ``Variance``/``StandardDeviation``.
The trailing tensor axis is the first geometry axis (numpy order is reversed), so that axis is
dropped from ``Origin``/``Spacing``/``Direction``.
"""
def __init__(self) -> None:
super().__init__()
def __call__(self, name: str, tensors: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if "Origin" in cache_attribute:
origin = cache_attribute.pop_np_array("Origin")
spacing = cache_attribute.pop_np_array("Spacing")
direction = cache_attribute.pop_np_array("Direction")
rank = len(origin)
cache_attribute["Origin"] = origin[1:]
cache_attribute["Spacing"] = spacing[1:]
cache_attribute["Direction"] = direction.reshape(rank, rank)[1:, 1:].flatten()
return torch.linalg.norm(tensors.float(), dim=-1)
[docs]
class Variance(Transform):
def __init__(self) -> None:
super().__init__()
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Variance across the leading member axis at each voxel -- no spatial neighbour.
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensors: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
# Keep the leading member axis in BOTH branches: the N>1 var(0) drops it and re-adds it via
# unsqueeze(0), so the single-member zeros must unsqueeze too or the output rank is off by one.
return (
tensors.float().var(0).unsqueeze(0) if tensors.shape[0] > 1 else torch.zeros_like(tensors[0]).unsqueeze(0)
)
[docs]
class SegmentationDisagreement(Transform):
def __init__(self, ignore_background: bool = False) -> None:
super().__init__()
self.ignore_background = ignore_background
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Per-voxel majority disagreement across the members. The global torch.unique only widens the
# label set with labels absent at a given voxel, which contribute zero counts there and never
# change that voxel's majority -- so the result is decided voxel by voxel.
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensors: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
# tensors shape: [N, ...] with N segmentations and integer labels per voxel
if tensors.shape[0] <= 1:
return torch.zeros_like(tensors[0], dtype=torch.float32).unsqueeze(0)
tensors = tensors.long()
if self.ignore_background:
valid = tensors != 0
else:
valid = torch.ones_like(tensors, dtype=torch.bool)
disagreement = torch.zeros_like(tensors[0], dtype=torch.float32)
# per-voxel disagreement = 1 - (frequency of majority label / number of valid segmentations)
unique_labels = torch.unique(tensors)
counts = []
for label in unique_labels:
counts.append(((tensors == label) & valid).sum(dim=0))
counts = torch.stack(counts, dim=0) # [L, ...]
max_count = counts.max(dim=0).values
valid_count = valid.sum(dim=0)
non_empty = valid_count > 0
disagreement[non_empty] = 1.0 - (max_count[non_empty].float() / valid_count[non_empty].float())
return disagreement.unsqueeze(0)
[docs]
class Percentage(Transform):
def __init__(self, baseline: float) -> None:
super().__init__()
self.baseline = baseline
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensors: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return tensors / self.baseline * 100.0
[docs]
class StandardDeviation(Transform):
def __init__(self) -> None:
super().__init__()
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Standard deviation across the leading member axis at each voxel -- no spatial neighbour.
return PatchLocality(LocalityKind.POINTWISE)
def __call__(self, name: str, tensors: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
return (
tensors.float().std(0).unsqueeze(0) if tensors.shape[0] > 1 else torch.zeros_like(tensors[0]).unsqueeze(0)
)
[docs]
class Statistics(Transform):
def __init__(self) -> None:
super().__init__()
def __call__(self, name: str, tensors: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
cache_attribute["ImageMin"] = tensors.float().min()
cache_attribute["ImageMax"] = tensors.float().max()
cache_attribute["ImageMean"] = tensors.float().mean()
cache_attribute["ImageStd"] = tensors.float().std()
return tensors
[docs]
class Crop(TransformInverse):
"""Crop a volume to the bounding box of its foreground.
The content-dependent box is computed once (``transform_shape``) and kept on the case as ``box``
margins; cropping is then the translation ``out[o] = volume[o + start]``, so a target patch reads
its shifted source region. Dropped voxels mean the stored volume's statistics are not the output's
(hence ``LocalityKind.CROP``).
"""
def __init__(self, inverse: bool = True) -> None:
super().__init__(inverse)
[docs]
def patch_locality(self, cache_attribute: Attribute) -> PatchLocality:
# Total: the box is a fact ``transform_shape`` puts on the case before the dispatcher reads any
# declaration, but a group carries only what its writer stored, and without it there is no
# translation to make -- only the read that would find one.
if "box" not in cache_attribute:
return PatchLocality(LocalityKind.WHOLE_VOLUME)
return PatchLocality(LocalityKind.CROP)
[docs]
def stream_region_source(
self,
target_slices: tuple[slice, ...],
source_spatial_shape: list[int],
cache_attribute: Attribute,
) -> list[slice]:
# Output index o holds source index o + start, so the region behind a target patch is that
# patch's own slices stepped forward by the box's near margin.
box = Crop._parse_box(cache_attribute["box"])
return [
slice(target.start + int(start), target.stop + int(start))
for target, (start, _) in zip(target_slices, box, strict=False)
]
[docs]
def write_stream_cache_attribute(self, cache_attribute: Attribute, source_spatial_shape: list[int]) -> None:
if "box" not in cache_attribute:
return
if not {"Origin", "Spacing", "Direction"} <= set(cache_attribute.keys()):
return
# The crop keeps the box's near corner, so the new origin is the physical point that corner
# already sat on: the old origin stepped along each axis by its own margin. A margin is in
# array order and the geometry is in (x, y, z), hence the reversed indexing.
box = Crop._parse_box(cache_attribute["box"])
origin = torch.tensor(cache_attribute.get_np_array("Origin"))
matrix = torch.tensor(cache_attribute.get_np_array("Direction").reshape((len(origin), len(origin))))
origin = torch.matmul(origin, matrix)
for dim in range(box.shape[0]):
origin[-dim - 1] += box[dim][0] * cache_attribute.get_np_array("Spacing")[-dim - 1]
cache_attribute["Origin"] = torch.matmul(origin, torch.inverse(matrix))
@staticmethod
def _parse_box(box_str: str) -> np.ndarray:
flat = np.fromstring(box_str.replace("[", " ").replace("]", " "), sep=" ", dtype=np.int64)
return flat.reshape(-1, 2)
def __call__(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if "box" not in cache_attribute:
return tensor
box = self._parse_box(cache_attribute["box"])
self.write_stream_cache_attribute(cache_attribute, list(tensor.shape[1:]))
# The box carries the FAR margin, so the stop it crops at is the one the extent in hand decides.
for i, ((_, b), s) in enumerate(zip(box, tensor.shape[1:], strict=False)):
box[i][1] = s - b
image = data_to_image(tensor, cache_attribute)
result = crop_with_mask(image, box)
data, _ = image_to_data(result)
return torch.from_numpy(data)
[docs]
def inverse(self, name: str, tensor: torch.Tensor, cache_attribute: Attribute) -> torch.Tensor:
if "box" not in cache_attribute:
return tensor
box = self._parse_box(cache_attribute.pop("box"))
cache_attribute.pop_np_array("Origin")
padding = []
for b in reversed(box):
padding.extend([b[0], b[1]])
result = F.pad(tensor.unsqueeze(0), tuple(padding), "replicate").squeeze(0)
return result