From cf074d69098adf37a38e026ead63107bb44f24c3 Mon Sep 17 00:00:00 2001 From: gbranaa4-hue Date: Sun, 6 Sep 2026 11:54:40 -0700 Subject: [PATCH 1/3] fix(networks): AffineHead follows the module's dtype, not just device self.grid is a plain attribute (torch.stack(...).to(dtype=torch.float)), never register_buffer'd. forward() re-derived only its device (self.grid.to(device=f.device)) every call, never its dtype. Casting a GlobalNet/LocalNet (both build on AffineHead) to half precision left self.grid at float32 while theta (from self.fc, whose weights did move) became float16 -- the very first half-precision forward call crashed: RuntimeError: expected scalar type Half but found Float at affine_transform's torch.einsum, which requires matching dtypes. Fix: re-derive dtype alongside device every forward call, mirroring the input f's dtype -- the same reference point self.fc(f...) already requires matching, so this is guaranteed consistent with theta's dtype by the time affine_transform(theta) - self.grid runs. Verified: half-precision GlobalNet forward no longer crashes (both 2D and via the full network), returns float16 output at the correct shape; float32 path is unaffected (byte-identical, since f.dtype was already torch.float in that case). Added test_half_precision_forward to test_globalnet.py; fails with the predicted RuntimeError on unpatched code, passes with the fix. Full existing GlobalNet/RegUNet test suites pass unchanged. Co-Authored-By: Claude Sonnet 5 Signed-off-by: tritsystem --- monai/networks/nets/regunet.py | 2 +- tests/networks/nets/test_globalnet.py | 14 ++++++++++++++ 2 files changed, 15 insertions(+), 1 deletion(-) diff --git a/monai/networks/nets/regunet.py b/monai/networks/nets/regunet.py index 4d6150ea1be..c8f9829bae2 100644 --- a/monai/networks/nets/regunet.py +++ b/monai/networks/nets/regunet.py @@ -295,7 +295,7 @@ def affine_transform(self, theta: torch.Tensor): def forward(self, x: list[torch.Tensor], image_size: list[int]) -> torch.Tensor: f = x[0] - self.grid = self.grid.to(device=f.device) + self.grid = self.grid.to(device=f.device, dtype=f.dtype) theta = self.fc(f.reshape(f.shape[0], -1)) if self.save_theta: self.theta = theta.detach() diff --git a/tests/networks/nets/test_globalnet.py b/tests/networks/nets/test_globalnet.py index ecb0243a1be..fcb0218aec1 100644 --- a/tests/networks/nets/test_globalnet.py +++ b/tests/networks/nets/test_globalnet.py @@ -98,6 +98,20 @@ def test_script(self, input_param, input_shape, _): test_data = torch.randn(input_shape) test_script_save(net, test_data) + @parameterized.expand(TEST_CASES_GLOBAL_NET) + def test_half_precision_forward(self, input_param, input_shape, expected_shape): + # AffineHead.grid was a plain float32 attribute, moved by device only + # (`self.grid.to(device=f.device)`) in forward(). Casting the network to half + # precision left `self.grid` at float32 while theta became float16, and + # `affine_transform`'s einsum raised `RuntimeError: expected scalar type Half + # but found Float` on the very first half-precision forward call. + net = GlobalNet(**input_param).half() + with eval_mode(net): + img = torch.randn(input_shape, dtype=torch.float16) + result = net(img) + self.assertEqual(result.dtype, torch.float16) + self.assertEqual(result.shape, expected_shape) + if __name__ == "__main__": unittest.main() From 9a258e98014384ebded10fd89e9136625fe03f5c Mon Sep 17 00:00:00 2001 From: gbranaa4-hue Date: Sun, 6 Sep 2026 11:54:40 -0700 Subject: [PATCH 2/3] fix(networks): AffineHead follows the module's dtype, not just device self.grid is a plain attribute (torch.stack(...).to(dtype=torch.float)), never register_buffer'd. forward() re-derived only its device (self.grid.to(device=f.device)) every call, never its dtype. Casting a GlobalNet/LocalNet (both build on AffineHead) to half precision left self.grid at float32 while theta (from self.fc, whose weights did move) became float16 -- the very first half-precision forward call crashed: RuntimeError: expected scalar type Half but found Float at affine_transform's torch.einsum, which requires matching dtypes. Fix: re-derive dtype alongside device every forward call, mirroring the input f's dtype -- the same reference point self.fc(f...) already requires matching, so this is guaranteed consistent with theta's dtype by the time affine_transform(theta) - self.grid runs. Verified: half-precision GlobalNet forward no longer crashes (both 2D and via the full network), returns float16 output at the correct shape; float32 path is unaffected (byte-identical, since f.dtype was already torch.float in that case). Added test_half_precision_forward to test_globalnet.py; fails with the predicted RuntimeError on unpatched code, passes with the fix. Full existing GlobalNet/RegUNet test suites pass unchanged. Co-Authored-By: Claude Sonnet 5 Signed-off-by: tritsystem --- monai/networks/nets/regunet.py | 2 +- tests/networks/nets/test_globalnet.py | 14 ++++++++++++++ 2 files changed, 15 insertions(+), 1 deletion(-) diff --git a/monai/networks/nets/regunet.py b/monai/networks/nets/regunet.py index 4d6150ea1be..c8f9829bae2 100644 --- a/monai/networks/nets/regunet.py +++ b/monai/networks/nets/regunet.py @@ -295,7 +295,7 @@ def affine_transform(self, theta: torch.Tensor): def forward(self, x: list[torch.Tensor], image_size: list[int]) -> torch.Tensor: f = x[0] - self.grid = self.grid.to(device=f.device) + self.grid = self.grid.to(device=f.device, dtype=f.dtype) theta = self.fc(f.reshape(f.shape[0], -1)) if self.save_theta: self.theta = theta.detach() diff --git a/tests/networks/nets/test_globalnet.py b/tests/networks/nets/test_globalnet.py index ecb0243a1be..fcb0218aec1 100644 --- a/tests/networks/nets/test_globalnet.py +++ b/tests/networks/nets/test_globalnet.py @@ -98,6 +98,20 @@ def test_script(self, input_param, input_shape, _): test_data = torch.randn(input_shape) test_script_save(net, test_data) + @parameterized.expand(TEST_CASES_GLOBAL_NET) + def test_half_precision_forward(self, input_param, input_shape, expected_shape): + # AffineHead.grid was a plain float32 attribute, moved by device only + # (`self.grid.to(device=f.device)`) in forward(). Casting the network to half + # precision left `self.grid` at float32 while theta became float16, and + # `affine_transform`'s einsum raised `RuntimeError: expected scalar type Half + # but found Float` on the very first half-precision forward call. + net = GlobalNet(**input_param).half() + with eval_mode(net): + img = torch.randn(input_shape, dtype=torch.float16) + result = net(img) + self.assertEqual(result.dtype, torch.float16) + self.assertEqual(result.shape, expected_shape) + if __name__ == "__main__": unittest.main() From 6012454ff3f1e2708fc5598bd754de3c46f6a4cf Mon Sep 17 00:00:00 2001 From: tritsystem Date: Mon, 21 Sep 2026 02:12:35 -0700 Subject: [PATCH 3/3] docs(tests): add Google-style docstrings requested in review Signed-off-by: tritsystem --- monai/networks/nets/regunet.py | 11 +++++++++++ tests/networks/nets/test_globalnet.py | 16 +++++++++++----- 2 files changed, 22 insertions(+), 5 deletions(-) diff --git a/monai/networks/nets/regunet.py b/monai/networks/nets/regunet.py index c8f9829bae2..9b659b21e11 100644 --- a/monai/networks/nets/regunet.py +++ b/monai/networks/nets/regunet.py @@ -294,6 +294,17 @@ def affine_transform(self, theta: torch.Tensor): return grid_warped def forward(self, x: list[torch.Tensor], image_size: list[int]) -> torch.Tensor: + """ + Predict an affine transform from the first feature map and return it as a displacement field. + + Args: + x: feature maps from the network; only the first, ``x[0]``, is used. The reference grid is moved to its + device and dtype, so a network cast to half precision stays in half precision. + image_size: spatial size of the input image; not used by this head. + + Returns: + The displacement field, the warped reference grid minus the reference grid. + """ f = x[0] self.grid = self.grid.to(device=f.device, dtype=f.dtype) theta = self.fc(f.reshape(f.shape[0], -1)) diff --git a/tests/networks/nets/test_globalnet.py b/tests/networks/nets/test_globalnet.py index fcb0218aec1..84e3d466dec 100644 --- a/tests/networks/nets/test_globalnet.py +++ b/tests/networks/nets/test_globalnet.py @@ -100,11 +100,17 @@ def test_script(self, input_param, input_shape, _): @parameterized.expand(TEST_CASES_GLOBAL_NET) def test_half_precision_forward(self, input_param, input_shape, expected_shape): - # AffineHead.grid was a plain float32 attribute, moved by device only - # (`self.grid.to(device=f.device)`) in forward(). Casting the network to half - # precision left `self.grid` at float32 while theta became float16, and - # `affine_transform`'s einsum raised `RuntimeError: expected scalar type Half - # but found Float` on the very first half-precision forward call. + """Check that a GlobalNet cast to half precision can run a forward pass and returns float16. + + `AffineHead.grid` used to be moved by device only, so casting the network left it at float32 while theta + became float16, and the einsum in `affine_transform` raised `RuntimeError: expected scalar type Half but + found Float` on the first half-precision forward call. + + Args: + input_param: keyword arguments used to build the GlobalNet. + input_shape: shape of the random float16 input. + expected_shape: shape the returned displacement field must have. + """ net = GlobalNet(**input_param).half() with eval_mode(net): img = torch.randn(input_shape, dtype=torch.float16)