From a555e95e1c3f7bb058de22842a6f2e76872e6eaf Mon Sep 17 00:00:00 2001 From: Richard Brown <33289025+rijobro@users.noreply.github.com> Date: Wed, 12 Jan 2022 14:05:44 +0000 Subject: [PATCH 1/2] avoid re-creation of objects Signed-off-by: Richard Brown <33289025+rijobro@users.noreply.github.com> --- monai/inferers/inferer.py | 3 ++- monai/transforms/utility/array.py | 6 ++++-- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/monai/inferers/inferer.py b/monai/inferers/inferer.py index c7b70e06ca..0b700d7c2e 100644 --- a/monai/inferers/inferer.py +++ b/monai/inferers/inferer.py @@ -30,7 +30,8 @@ class Inferer(ABC): Example code:: device = torch.device("cuda:0") - data = ToTensor()(LoadImage()(filename=img_path)).to(device) + transform = Compose([ToTensor(), LoadImage(image_only=True)]) + data = transform(img_path).to(device) model = UNet(...).to(device) inferer = SlidingWindowInferer(...) diff --git a/monai/transforms/utility/array.py b/monai/transforms/utility/array.py index a107cf1cb1..0053b3c904 100644 --- a/monai/transforms/utility/array.py +++ b/monai/transforms/utility/array.py @@ -202,6 +202,8 @@ def __init__(self, strict_check: bool = True): strict_check: whether to raise an error when the meta information is insufficient. """ self.strict_check = strict_check + self.add_channel = AddChannel() + self.as_channel_first = AsChannelFirst() def __call__(self, img: NdarrayOrTensor, meta_dict: Optional[Mapping] = None) -> NdarrayOrTensor: """ @@ -223,8 +225,8 @@ def __call__(self, img: NdarrayOrTensor, meta_dict: Optional[Mapping] = None) -> warnings.warn(msg) return img if channel_dim == "no_channel": - return AddChannel()(img) - return AsChannelFirst(channel_dim=channel_dim)(img) + return self.add_channel(img) + return self.as_channel_first(channel_dim=channel_dim)(img) class RepeatChannel(Transform): From f04894488f180368adf4283a9a8c39ca5a7aad23 Mon Sep 17 00:00:00 2001 From: Richard Brown <33289025+rijobro@users.noreply.github.com> Date: Wed, 12 Jan 2022 14:18:37 +0000 Subject: [PATCH 2/2] AsChannelFirst needs to be created during __call__ Signed-off-by: Richard Brown <33289025+rijobro@users.noreply.github.com> --- monai/transforms/utility/array.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/monai/transforms/utility/array.py b/monai/transforms/utility/array.py index 0053b3c904..664433270b 100644 --- a/monai/transforms/utility/array.py +++ b/monai/transforms/utility/array.py @@ -203,7 +203,6 @@ def __init__(self, strict_check: bool = True): """ self.strict_check = strict_check self.add_channel = AddChannel() - self.as_channel_first = AsChannelFirst() def __call__(self, img: NdarrayOrTensor, meta_dict: Optional[Mapping] = None) -> NdarrayOrTensor: """ @@ -226,7 +225,7 @@ def __call__(self, img: NdarrayOrTensor, meta_dict: Optional[Mapping] = None) -> return img if channel_dim == "no_channel": return self.add_channel(img) - return self.as_channel_first(channel_dim=channel_dim)(img) + return AsChannelFirst(channel_dim=channel_dim)(img) class RepeatChannel(Transform):