Skip to content

Commit 26326b5

Browse files
aymuos15ericspod
andauthored
Add explicit spatial_ndim tracking to MetaTensor (#8765)
## Summary Fixes #6397 - Adds a `_spatial_ndim: int` attribute to `MetaTensor` that explicitly tracks the number of spatial dimensions, preventing dimension-mismatch crashes when `einops.rearrange()` or other reshape operations change `ndim` - The attribute propagates through `copy_meta_from` (via `__dict__` copy) and is preserved through arbitrary torch operations - Updates transforms (`Resize`, `Rotate`, `Zoom`, `Flip`, `Affine`, `SplitDim`, `AddCoordinateChannels`, etc.) and lazy resampling to use `spatial_ndim` instead of hardcoded 3 ### Key design decisions - **Constructor**: `spatial_ndim = min(affine.shape[-1] - 1, ndim - 1)` — clamped by actual tensor dims - **Affine setter**: `spatial_ndim = affine.shape[-1] - 1` — no clamping (user is explicit) - **`peek_pending_affine`**: uses affine's inner matrix shape (fixes batched `(1,4,4)` case) - **`spatial_resample`**: `min(spatial_ndim, ndim - 1, 3)` — adds ndim-1 constraint as safety net ### Files changed (16) - `monai/data/meta_obj.py`, `meta_tensor.py`, `utils.py`, `__init__.py` — core MetaTensor changes - `monai/transforms/` — spatial, croppad, intensity, inverse, lazy, post, utility transforms updated - `tests/data/meta_tensor/test_spatial_ndim.py` — 18 new tests - Existing test files updated with `spatial_ndim` assertions ## Test plan - [x] 18 new unit tests for `spatial_ndim` property (construction, affine sync, propagation, einops reshape, transforms) - [x] Existing MetaTensor tests pass (162 tests) - [x] SqueezeDim and SplitDim tests pass with new assertions - [x] Total: 216 tests verified passing --------- Signed-off-by: Soumya Snigdha Kundu <soumya_snigdha.kundu@kcl.ac.uk> Co-authored-by: Eric Kerfoot <17726042+ericspod@users.noreply.github.com>
1 parent 31419ab commit 26326b5

16 files changed

Lines changed: 414 additions & 57 deletions

File tree

monai/data/__init__.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -71,7 +71,7 @@
7171
monai_to_itk_ddf,
7272
)
7373
from .meta_obj import MetaObj, get_track_meta, set_track_meta
74-
from .meta_tensor import MetaTensor
74+
from .meta_tensor import MetaTensor, get_spatial_ndim
7575
from .samplers import DistributedSampler, DistributedWeightedRandomSampler
7676
from .synthetic import create_test_image_2d, create_test_image_3d
7777
from .test_time_augmentation import TestTimeAugmentation

monai/data/meta_obj.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,9 @@
2424

2525
_TRACK_META = True
2626

27+
# Default number of spatial dimensions for medical imaging (3D volumetric data)
28+
_DEFAULT_SPATIAL_NDIM = 3
29+
2730
__all__ = ["get_track_meta", "set_track_meta", "MetaObj"]
2831

2932

@@ -84,6 +87,7 @@ def __init__(self) -> None:
8487
self._applied_operations: list = MetaObj.get_default_applied_operations()
8588
self._pending_operations: list = MetaObj.get_default_applied_operations() # the same default as applied_ops
8689
self._is_batch: bool = False
90+
self._spatial_ndim: int = 3 # default: 3 spatial dimensions
8791

8892
@staticmethod
8993
def flatten_meta_objs(*args: Iterable):

monai/data/meta_tensor.py

Lines changed: 94 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -13,22 +13,60 @@
1313

1414
import functools
1515
import warnings
16-
from collections.abc import Sequence
16+
from collections.abc import Mapping, Sequence
1717
from copy import deepcopy
18+
from numbers import Integral
1819
from typing import Any
1920

2021
import numpy as np
2122
import torch
2223

2324
import monai
24-
from monai.config.type_definitions import NdarrayTensor
25-
from monai.data.meta_obj import MetaObj, get_track_meta
26-
from monai.data.utils import affine_to_spacing, decollate_batch, list_data_collate, remove_extra_metadata
25+
from monai.config.type_definitions import NdarrayOrTensor, NdarrayTensor
26+
from monai.data.meta_obj import _DEFAULT_SPATIAL_NDIM, MetaObj, get_track_meta
27+
from monai.data.utils import affine_to_spacing, decollate_batch, is_no_channel, list_data_collate, remove_extra_metadata
2728
from monai.utils import look_up_option
2829
from monai.utils.enums import LazyAttr, MetaKeys, PostFix, SpaceKeys
2930
from monai.utils.type_conversion import convert_data_type, convert_to_dst_type, convert_to_numpy, convert_to_tensor
3031

31-
__all__ = ["MetaTensor"]
32+
__all__ = ["MetaTensor", "get_spatial_ndim"]
33+
34+
35+
def _normalize_spatial_ndim(spatial_ndim: int, tensor_ndim: int, no_channel: bool = False) -> int:
36+
"""Clamp spatial dims to a valid range for the current tensor shape."""
37+
limit = max(int(tensor_ndim), 1) if no_channel else max(int(tensor_ndim) - 1, 1)
38+
return max(1, min(int(spatial_ndim), limit))
39+
40+
41+
def _has_explicit_no_channel(meta: Mapping | None) -> bool:
42+
return (
43+
isinstance(meta, Mapping)
44+
and MetaKeys.ORIGINAL_CHANNEL_DIM in meta
45+
and is_no_channel(meta[MetaKeys.ORIGINAL_CHANNEL_DIM])
46+
)
47+
48+
49+
def get_spatial_ndim(img: NdarrayOrTensor) -> int:
50+
"""Return the number of spatial dimensions assuming channel-first layout.
51+
52+
Uses ``MetaTensor.spatial_ndim`` when available, otherwise falls back to
53+
``img.ndim - 1``. Always assumes channel-first (``no_channel=False``)
54+
because callers run after ``EnsureChannelFirst`` has already added one.
55+
"""
56+
if isinstance(img, MetaTensor):
57+
return _normalize_spatial_ndim(img.spatial_ndim, img.ndim)
58+
return img.ndim - 1
59+
60+
61+
def _is_batch_only_index(index: Any) -> bool:
62+
"""True when indexing pattern selects only the batch axis (e.g., ``x[0]`` or ``x[0, ...]``)."""
63+
if isinstance(index, (int, np.integer)):
64+
return True
65+
if not isinstance(index, Sequence) or not index:
66+
return False
67+
if not isinstance(index[0], (int, np.integer)):
68+
return False
69+
return all(i in (slice(None, None, None), Ellipsis, None) for i in index[1:])
3270

3371

3472
@functools.lru_cache(None)
@@ -111,6 +149,7 @@ def __new__(
111149
meta: dict | None = None,
112150
applied_operations: list | None = None,
113151
*args,
152+
spatial_ndim: int | None = None,
114153
**kwargs,
115154
) -> MetaTensor:
116155
_kwargs = {"device": kwargs.pop("device", None), "dtype": kwargs.pop("dtype", None)} if kwargs else {}
@@ -123,6 +162,7 @@ def __init__(
123162
meta: dict | None = None,
124163
applied_operations: list | None = None,
125164
*_args,
165+
spatial_ndim: int | None = None,
126166
**_kwargs,
127167
) -> None:
128168
"""
@@ -134,6 +174,8 @@ def __init__(
134174
the list is typically maintained by `monai.transforms.TraceableTransform`.
135175
See also: :py:class:`monai.transforms.TraceableTransform`
136176
_args: additional args (currently not in use in this constructor).
177+
spatial_ndim: optional number of spatial dimensions. If ``None``, derived
178+
from the affine matrix clamped by the tensor shape.
137179
_kwargs: additional kwargs (currently not in use in this constructor).
138180
139181
Note:
@@ -158,6 +200,14 @@ def __init__(
158200
self.affine = self.meta[MetaKeys.AFFINE]
159201
else:
160202
self.affine = self.get_default_affine()
203+
# Initialize spatial_ndim from affine matrix (source of truth), clamped by tensor shape.
204+
# This cached value is kept in sync via the affine setter for hot-path performance.
205+
no_channel = _has_explicit_no_channel(self.meta)
206+
if spatial_ndim is not None:
207+
self.spatial_ndim = _normalize_spatial_ndim(spatial_ndim, self.ndim, no_channel=no_channel)
208+
elif self.affine.ndim == 2:
209+
self.spatial_ndim = _normalize_spatial_ndim(self.affine.shape[-1] - 1, self.ndim, no_channel=no_channel)
210+
161211
# applied_operations
162212
if applied_operations is not None:
163213
self.applied_operations = applied_operations
@@ -237,6 +287,7 @@ def _handle_batched(cls, ret, idx, metas, func, args, kwargs):
237287
if func == torch.Tensor.__getitem__:
238288
if idx > 0 or len(args) < 2 or len(args[0]) < 1:
239289
return ret
290+
full_idx = args[1]
240291
batch_idx = args[1][0] if isinstance(args[1], Sequence) else args[1]
241292
# if using e.g., `batch[:, -1]` or `batch[..., -1]`, then the
242293
# first element will be `slice(None, None, None)` and `Ellipsis`,
@@ -258,6 +309,8 @@ def _handle_batched(cls, ret, idx, metas, func, args, kwargs):
258309
ret_meta.is_batch = False
259310
if hasattr(ret_meta, "__dict__"):
260311
ret.__dict__ = ret_meta.__dict__.copy()
312+
if _is_batch_only_index(full_idx):
313+
ret.spatial_ndim = _normalize_spatial_ndim(ret.spatial_ndim, ret.ndim)
261314
# `unbind` is used for `next(iter(batch))`. Also for `decollate_batch`.
262315
# But we only want to split the batch if the `unbind` is along the 0th dimension.
263316
elif func == torch.Tensor.unbind:
@@ -467,15 +520,42 @@ def affine(self) -> torch.Tensor:
467520

468521
@affine.setter
469522
def affine(self, d: NdarrayTensor) -> None:
470-
"""Set the affine."""
471-
self.meta[MetaKeys.AFFINE] = torch.as_tensor(d, device=torch.device("cpu"), dtype=torch.float64)
523+
"""Set the affine.
524+
525+
When setting a non-batched affine matrix, automatically synchronizes the cached
526+
spatial_ndim attribute to maintain consistency between the affine matrix (source of truth)
527+
and the cached spatial dimension count.
528+
"""
529+
a = torch.as_tensor(d, device=torch.device("cpu"), dtype=torch.float64)
530+
self.meta[MetaKeys.AFFINE] = a
531+
if a.ndim == 2: # non-batched: sync spatial_ndim from affine (source of truth)
532+
no_channel = _has_explicit_no_channel(self.meta)
533+
self.spatial_ndim = _normalize_spatial_ndim(a.shape[-1] - 1, self.ndim, no_channel=no_channel)
534+
535+
@property
536+
def spatial_ndim(self) -> int:
537+
"""Get the number of spatial dimensions.
538+
539+
This value is cached for hot-path performance and is kept in sync with the affine matrix
540+
via the affine setter. The affine matrix is the source of truth for spatial dimensions.
541+
"""
542+
return getattr(self, "_spatial_ndim", _DEFAULT_SPATIAL_NDIM)
543+
544+
@spatial_ndim.setter
545+
def spatial_ndim(self, val: int) -> None:
546+
"""Set the number of spatial dimensions."""
547+
if not isinstance(val, Integral):
548+
raise TypeError(f"'val' must be an numbers.Integral type; got {type(val)}.")
549+
if val < 1:
550+
raise ValueError(f"spatial_ndim must be >= 1, got {val}")
551+
self._spatial_ndim = int(val)
472552

473553
@property
474554
def pixdim(self):
475555
"""Get the spacing"""
476556
if self.is_batch:
477-
return [affine_to_spacing(a) for a in self.affine]
478-
return affine_to_spacing(self.affine)
557+
return [affine_to_spacing(a, r=self.spatial_ndim) for a in self.affine]
558+
return affine_to_spacing(self.affine, r=self.spatial_ndim)
479559

480560
def peek_pending_shape(self):
481561
"""
@@ -490,7 +570,7 @@ def peek_pending_shape(self):
490570

491571
def peek_pending_affine(self):
492572
res = self.affine
493-
r = len(res) - 1
573+
r = res.shape[-1] - 1 if res.ndim >= 2 else self.spatial_ndim
494574
if r not in (2, 3):
495575
warnings.warn(f"Only 2d and 3d affine are supported, got {r}d input.")
496576
for p in self.pending_operations:
@@ -503,8 +583,10 @@ def peek_pending_affine(self):
503583
return res
504584

505585
def peek_pending_rank(self):
506-
a = self.pending_operations[-1].get(LazyAttr.AFFINE, None) if self.pending_operations else self.affine
507-
return 1 if a is None else int(max(1, len(a) - 1))
586+
if self.pending_operations:
587+
a = self.pending_operations[-1].get(LazyAttr.AFFINE, None)
588+
return 1 if a is None else int(max(1, len(a) - 1))
589+
return self.spatial_ndim
508590

509591
def new_empty(self, size, dtype=None, device=None, requires_grad=False): # type: ignore[override]
510592
"""

monai/data/utils.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -31,7 +31,7 @@
3131
from torch.utils.data._utils.collate import default_collate
3232

3333
from monai.config.type_definitions import NdarrayOrTensor, NdarrayTensor, PathLike
34-
from monai.data.meta_obj import MetaObj
34+
from monai.data.meta_obj import _DEFAULT_SPATIAL_NDIM, MetaObj
3535
from monai.utils import (
3636
MAX_SEED,
3737
BlendMode,
@@ -432,6 +432,9 @@ def collate_meta_tensor_fn(batch, *, collate_fn_map=None):
432432
collated.meta = default_collate(meta_dicts)
433433
collated.applied_operations = [i.applied_operations or TraceKeys.NONE for i in batch]
434434
collated.is_batch = True
435+
collated.spatial_ndim = min(
436+
min(getattr(t, "spatial_ndim", _DEFAULT_SPATIAL_NDIM) for t in batch), max(collated.ndim - 1, 1)
437+
)
435438
return collated
436439

437440

monai/transforms/croppad/functional.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222

2323
from monai.config.type_definitions import NdarrayTensor
2424
from monai.data.meta_obj import get_track_meta
25-
from monai.data.meta_tensor import MetaTensor
25+
from monai.data.meta_tensor import MetaTensor, get_spatial_ndim
2626
from monai.data.utils import to_affine_nd
2727
from monai.transforms.inverse import TraceableTransform
2828
from monai.transforms.utils import convert_pad_mode, create_translate
@@ -132,7 +132,7 @@ def crop_or_pad_nd(img: torch.Tensor, translation_mat, spatial_size: tuple[int,
132132
mode: the padding mode.
133133
kwargs: other arguments for the `np.pad` or `torch.pad` function.
134134
"""
135-
ndim = len(img.shape) - 1
135+
ndim = get_spatial_ndim(img)
136136
matrix_np = np.round(to_affine_nd(ndim, convert_to_numpy(translation_mat, wrap_sequence=True).copy()))
137137
matrix_np = to_affine_nd(len(spatial_size), matrix_np)
138138
cc = np.asarray(np.meshgrid(*[[0.5, x - 0.5] for x in spatial_size], indexing="ij"))

monai/transforms/intensity/array.py

Lines changed: 8 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626
from monai.config import DtypeLike
2727
from monai.config.type_definitions import NdarrayOrTensor, NdarrayTensor
2828
from monai.data.meta_obj import get_track_meta
29+
from monai.data.meta_tensor import get_spatial_ndim
2930
from monai.data.ultrasound_confidence_map import UltrasoundConfidenceMap
3031
from monai.data.utils import get_random_patch, get_valid_patch_size
3132
from monai.networks.layers import GaussianFilter, HilbertTransform, MedianFilter, SavitzkyGolayFilter
@@ -1603,7 +1604,7 @@ def __init__(self, radius: Sequence[int] | int = 1) -> None:
16031604
def __call__(self, img: NdarrayTensor) -> NdarrayTensor:
16041605
img = convert_to_tensor(img, track_meta=get_track_meta())
16051606
img_t, *_ = convert_data_type(img, torch.Tensor, dtype=torch.float)
1606-
spatial_dims = img_t.ndim - 1
1607+
spatial_dims = get_spatial_ndim(img)
16071608
r = ensure_tuple_rep(self.radius, spatial_dims)
16081609
median_filter_instance = MedianFilter(r, spatial_dims=spatial_dims)
16091610
out_t: torch.Tensor = median_filter_instance(img_t)
@@ -1639,7 +1640,7 @@ def __call__(self, img: NdarrayTensor) -> NdarrayTensor:
16391640
sigma = [torch.as_tensor(s, device=img_t.device) for s in self.sigma]
16401641
else:
16411642
sigma = torch.as_tensor(self.sigma, device=img_t.device)
1642-
gaussian_filter = GaussianFilter(img_t.ndim - 1, sigma, approx=self.approx)
1643+
gaussian_filter = GaussianFilter(get_spatial_ndim(img), sigma, approx=self.approx)
16431644
out_t: torch.Tensor = gaussian_filter(img_t.unsqueeze(0)).squeeze(0)
16441645
out, *_ = convert_to_dst_type(out_t, dst=img, dtype=out_t.dtype)
16451646

@@ -1696,7 +1697,7 @@ def __call__(self, img: NdarrayOrTensor, randomize: bool = True) -> NdarrayOrTen
16961697
if not self._do_transform:
16971698
return img
16981699

1699-
sigma = ensure_tuple_size(vals=(self.x, self.y, self.z), dim=img.ndim - 1)
1700+
sigma = ensure_tuple_size(vals=(self.x, self.y, self.z), dim=get_spatial_ndim(img))
17001701
return GaussianSmooth(sigma=sigma, approx=self.approx)(img)
17011702

17021703

@@ -1746,7 +1747,7 @@ def __call__(self, img: NdarrayTensor) -> NdarrayTensor:
17461747
img_t, *_ = convert_data_type(img, torch.Tensor, dtype=torch.float32)
17471748

17481749
gf1, gf2 = (
1749-
GaussianFilter(img_t.ndim - 1, sigma, approx=self.approx).to(img_t.device)
1750+
GaussianFilter(get_spatial_ndim(img), sigma, approx=self.approx).to(img_t.device)
17501751
for sigma in (self.sigma1, self.sigma2)
17511752
)
17521753
blurred_f = gf1(img_t.unsqueeze(0))
@@ -1834,8 +1835,9 @@ def __call__(self, img: NdarrayOrTensor, randomize: bool = True) -> NdarrayOrTen
18341835

18351836
if self.x2 is None or self.y2 is None or self.z2 is None or self.a is None:
18361837
raise RuntimeError("please call the `randomize()` function first.")
1837-
sigma1 = ensure_tuple_size(vals=(self.x1, self.y1, self.z1), dim=img.ndim - 1)
1838-
sigma2 = ensure_tuple_size(vals=(self.x2, self.y2, self.z2), dim=img.ndim - 1)
1838+
_sp = get_spatial_ndim(img)
1839+
sigma1 = ensure_tuple_size(vals=(self.x1, self.y1, self.z1), dim=_sp)
1840+
sigma2 = ensure_tuple_size(vals=(self.x2, self.y2, self.z2), dim=_sp)
18391841
return GaussianSharpen(sigma1=sigma1, sigma2=sigma2, alpha=self.a, approx=self.approx)(img)
18401842

18411843

monai/transforms/inverse.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -215,7 +215,7 @@ def track_transform_meta(
215215
orig_affine = data_t.peek_pending_affine()
216216
orig_affine = convert_to_dst_type(orig_affine, affine, dtype=torch.float64)[0]
217217
try:
218-
affine = orig_affine @ to_affine_nd(len(orig_affine) - 1, affine, dtype=torch.float64)
218+
affine = orig_affine @ to_affine_nd(orig_affine.shape[-1] - 1, affine, dtype=torch.float64)
219219
except RuntimeError as e:
220220
if orig_affine.ndim > 2:
221221
if data_t.is_batch:

monai/transforms/lazy/functional.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -257,9 +257,11 @@ def apply_pending(data: torch.Tensor | MetaTensor, pending: list | None = None,
257257
if not pending:
258258
return data, []
259259

260+
_rank = data.spatial_ndim if isinstance(data, MetaTensor) else 3
261+
260262
cumulative_xform = affine_from_pending(pending[0])
261-
if cumulative_xform.shape[0] == 3:
262-
cumulative_xform = to_affine_nd(3, cumulative_xform)
263+
if cumulative_xform.shape[0] < _rank + 1:
264+
cumulative_xform = to_affine_nd(_rank, cumulative_xform)
263265

264266
cur_kwargs = kwargs_from_pending(pending[0])
265267
override_kwargs: dict[str, Any] = {}
@@ -284,8 +286,8 @@ def apply_pending(data: torch.Tensor | MetaTensor, pending: list | None = None,
284286
data = resample(data.to(device), cumulative_xform, _cur_kwargs)
285287

286288
next_matrix = affine_from_pending(p)
287-
if next_matrix.shape[0] == 3:
288-
next_matrix = to_affine_nd(3, next_matrix)
289+
if next_matrix.shape[0] < _rank + 1:
290+
next_matrix = to_affine_nd(_rank, next_matrix)
289291

290292
cumulative_xform = combine_transforms(cumulative_xform, next_matrix)
291293
cur_kwargs.update(new_kwargs)

monai/transforms/post/array.py

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@
2323

2424
from monai.config.type_definitions import NdarrayOrTensor
2525
from monai.data.meta_obj import get_track_meta
26-
from monai.data.meta_tensor import MetaTensor
26+
from monai.data.meta_tensor import MetaTensor, get_spatial_ndim
2727
from monai.networks import one_hot
2828
from monai.networks.layers import GaussianFilter, apply_filter, separable_filtering
2929
from monai.transforms.inverse import InvertibleTransform
@@ -624,7 +624,11 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
624624
"""
625625
img = convert_to_tensor(img, track_meta=get_track_meta())
626626
img_: torch.Tensor = convert_to_tensor(img, track_meta=False)
627-
spatial_dims = len(img_.shape) - 1
627+
spatial_dims = get_spatial_ndim(img)
628+
# Validate actual tensor shape against tracked spatial_ndim
629+
actual_spatial = img_.ndim - 1 # channel-first layout
630+
if actual_spatial != spatial_dims:
631+
spatial_dims = actual_spatial
628632
img_ = img_.unsqueeze(0) # adds a batch dim
629633
if spatial_dims == 2:
630634
kernel = torch.tensor([[-1, -1, -1], [-1, 8, -1], [-1, -1, -1]], dtype=torch.float32)
@@ -1104,7 +1108,7 @@ def __call__(self, image: NdarrayOrTensor) -> torch.Tensor:
11041108
image_tensor = convert_to_tensor(image, track_meta=get_track_meta())
11051109

11061110
# Check/set spatial axes
1107-
n_spatial_dims = image_tensor.ndim - 1 # excluding the channel dimension
1111+
n_spatial_dims = get_spatial_ndim(image_tensor)
11081112
valid_spatial_axes = list(range(n_spatial_dims)) + list(range(-n_spatial_dims, 0))
11091113

11101114
# Check gradient axes to be valid

0 commit comments

Comments
 (0)