diff --git a/monai/networks/nets/regunet.py b/monai/networks/nets/regunet.py index 4d6150ea1b..9b659b21e1 100644 --- a/monai/networks/nets/regunet.py +++ b/monai/networks/nets/regunet.py @@ -294,8 +294,19 @@ 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) + 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 ecb0243a1b..84e3d466de 100644 --- a/tests/networks/nets/test_globalnet.py +++ b/tests/networks/nets/test_globalnet.py @@ -98,6 +98,26 @@ 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): + """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) + result = net(img) + self.assertEqual(result.dtype, torch.float16) + self.assertEqual(result.shape, expected_shape) + if __name__ == "__main__": unittest.main()