Skip to content

Commit c575f22

Browse files
committed
perf: skip redundant convert_to_tensor for MetaTensor strip
Several transforms invoke ``convert_to_tensor`` twice on the same input: once with ``track_meta=get_track_meta()`` to keep the MetaTensor wrapper, then again with ``track_meta=False`` to obtain a plain-tensor view for the actual computation. The second call routes through ``_convert_tensor(data).to(dtype=..., device=..., memory_format=contiguous_format)`` (``monai/utils/type_conversion.py``) and is functionally equivalent to calling ``MetaTensor.as_tensor()`` for the purpose intended here, since the first call already enforced ``contiguous_format``. Replace the second call with an explicit ``as_tensor()`` strip: img_t = img.as_tensor() if isinstance(img, MetaTensor) else img This avoids a redundant ``convert_to_tensor`` dispatch and a no-op ``Tensor.to(memory_format=contiguous_format)`` per ``__call__``. Sites touched (all in transform ``__call__``, so per-sample per-epoch): * ``KeepLargestConnectedComponent`` (post/array.py) * ``RemoveSmallObjects`` (post/array.py) * ``LabelToContour`` (post/array.py) * ``ScaleIntensity`` (intensity/array.py) * ``ScaleIntensityFixedMean`` (intensity/array.py) * ``ClipIntensityPercentiles`` (intensity/array.py) * ``NormalizeIntensity`` (intensity/array.py) * ``SavitzkyGolaySmooth`` (intensity/array.py) * ``RandHistogramShift`` (intensity/array.py) * ``KSpaceSpikeNoise`` (intensity/array.py) * ``IntensityRemap`` (intensity/array.py) No behavioral change: the resulting tensor is the same underlying data (``as_tensor()`` returns the wrapped ``torch.Tensor`` without copying). Signed-off-by: Soumya Snigdha Kundu <soumya_snigdha.kundu@kcl.ac.uk>
1 parent a4d1a0d commit c575f22

2 files changed

Lines changed: 13 additions & 11 deletions

File tree

monai/transforms/intensity/array.py

Lines changed: 9 additions & 8 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 MetaTensor
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
@@ -483,7 +484,7 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
483484
484485
"""
485486
img = convert_to_tensor(img, track_meta=get_track_meta())
486-
img_t = convert_to_tensor(img, track_meta=False)
487+
img_t = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
487488
ret: NdarrayOrTensor
488489
if self.minv is not None or self.maxv is not None:
489490
if self.channel_wise:
@@ -542,7 +543,7 @@ def __call__(self, img: NdarrayOrTensor, factor=None) -> NdarrayOrTensor:
542543
factor = factor if factor is not None else self.factor
543544

544545
img = convert_to_tensor(img, track_meta=get_track_meta())
545-
img_t = convert_to_tensor(img, track_meta=False)
546+
img_t = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
546547
ret: NdarrayOrTensor
547548
if self.channel_wise:
548549
out = []
@@ -1168,7 +1169,7 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
11681169
Apply the transform to `img`.
11691170
"""
11701171
img = convert_to_tensor(img, track_meta=get_track_meta())
1171-
img_t = convert_to_tensor(img, track_meta=False)
1172+
img_t = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
11721173
if self.channel_wise:
11731174
img_t = torch.stack([self._clip(img=d) for d in img_t]) # type: ignore
11741175
else:
@@ -1433,7 +1434,7 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
14331434
Apply the transform to `img`.
14341435
"""
14351436
img = convert_to_tensor(img, track_meta=get_track_meta())
1436-
img_t = convert_to_tensor(img, track_meta=False)
1437+
img_t = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
14371438
if self.channel_wise:
14381439
img_t = torch.stack([self._normalize(img=d) for d in img_t]) # type: ignore
14391440
else:
@@ -1530,7 +1531,7 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
15301531
15311532
"""
15321533
img = convert_to_tensor(img, track_meta=get_track_meta())
1533-
self.img_t = convert_to_tensor(img, track_meta=False)
1534+
self.img_t = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
15341535

15351536
# add one to transform axis because a batch axis will be added at dimension 0
15361537
savgol_filter = SavitzkyGolayFilter(self.window_length, self.order, self.axis + 1, self.mode)
@@ -1907,7 +1908,7 @@ def __call__(self, img: NdarrayOrTensor, randomize: bool = True) -> NdarrayOrTen
19071908

19081909
if self.reference_control_points is None or self.floating_control_points is None:
19091910
raise RuntimeError("please call the `randomize()` function first.")
1910-
img_t = convert_to_tensor(img, track_meta=False)
1911+
img_t = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
19111912
img_min, img_max = img_t.min(), img_t.max()
19121913
if img_min == img_max:
19131914
warn(
@@ -1952,7 +1953,7 @@ def __init__(self, alpha: float = 0.1) -> None:
19521953

19531954
def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
19541955
img = convert_to_tensor(img, track_meta=get_track_meta())
1955-
img_t = convert_to_tensor(img, track_meta=False)
1956+
img_t = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
19561957
n_dims = len(img_t.shape[1:])
19571958

19581959
# FT
@@ -2604,7 +2605,7 @@ def __call__(self, img: torch.Tensor) -> torch.Tensor:
26042605
img: image to remap.
26052606
"""
26062607
img = convert_to_tensor(img, track_meta=get_track_meta())
2607-
img_ = convert_to_tensor(img, track_meta=False)
2608+
img_ = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
26082609
# sample noise
26092610
vals_to_sample = torch.unique(img_).tolist()
26102611
noise = torch.from_numpy(self.R.choice(vals_to_sample, len(vals_to_sample) - 1 + self.kernel_size))

monai/transforms/post/array.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616

1717
import warnings
1818
from collections.abc import Callable, Iterable, Sequence
19+
from typing import cast
1920

2021
import numpy as np
2122
import torch
@@ -338,7 +339,7 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
338339
else:
339340
applied_labels = tuple(get_unique_labels(img, is_onehot, discard=0))
340341
img = convert_to_tensor(img, track_meta=get_track_meta())
341-
img_: torch.Tensor = convert_to_tensor(img, track_meta=False)
342+
img_: torch.Tensor = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
342343
if self.independent:
343344
for i in applied_labels:
344345
foreground = img_[i] > 0 if is_onehot else img_[0] == i
@@ -497,7 +498,7 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
497498

498499
if isinstance(img, torch.Tensor):
499500
img = convert_to_tensor(img, track_meta=get_track_meta())
500-
img_ = convert_to_tensor(img, track_meta=False)
501+
img_ = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
501502
if hasattr(torch, "isin"): # `isin` is new in torch 1.10.0
502503
appl_lbls = torch.as_tensor(self.applied_labels, device=img_.device)
503504
out = torch.where(torch.isin(img_, appl_lbls), img_, torch.tensor(0.0).to(img_))
@@ -623,7 +624,7 @@ def __call__(self, img: NdarrayOrTensor) -> NdarrayOrTensor:
623624
624625
"""
625626
img = convert_to_tensor(img, track_meta=get_track_meta())
626-
img_: torch.Tensor = convert_to_tensor(img, track_meta=False)
627+
img_: torch.Tensor = img.as_tensor() if isinstance(img, MetaTensor) else cast(torch.Tensor, img)
627628
spatial_dims = len(img_.shape) - 1
628629
img_ = img_.unsqueeze(0) # adds a batch dim
629630
if spatial_dims == 2:

0 commit comments

Comments
 (0)