diff --git a/monai/networks/blocks/warp.py b/monai/networks/blocks/warp.py index 1878662916..cd75bacaad 100644 --- a/monai/networks/blocks/warp.py +++ b/monai/networks/blocks/warp.py @@ -109,13 +109,27 @@ def __init__(self, mode=GridSampleMode.BILINEAR.value, padding_mode=GridSamplePa self._padding_mode = self._padding_mode_native self.ref_grid = None + self._ref_grid_params: tuple[bool, int | None] | None = None self.jitter = jitter def get_reference_grid(self, ddf: torch.Tensor, jitter: bool = False, seed: int = 0) -> torch.Tensor: + """Return a reference grid matching the displacement field and generation parameters. + + Args: + ddf: Dense displacement field defining the grid shape, device, and dtype. + jitter: Whether to add deterministic random offsets to the grid. + seed: Random seed used when ``jitter`` is enabled. + + Returns: + The cached or newly generated reference grid. + """ + ref_grid_params = (jitter, seed if jitter else None) if ( self.ref_grid is not None - and self.ref_grid.shape[0] == ddf.shape[0] - and self.ref_grid.shape[1:] == ddf.shape[2:] + and self.ref_grid.shape == ddf.shape + and self.ref_grid.device == ddf.device + and self.ref_grid.dtype == ddf.dtype + and self._ref_grid_params == ref_grid_params ): return self.ref_grid # type: ignore mesh_points = [torch.arange(0, dim) for dim in ddf.shape[2:]] @@ -128,6 +142,7 @@ def get_reference_grid(self, ddf: torch.Tensor, jitter: bool = False, seed: int torch.random.manual_seed(seed) grid += torch.rand_like(grid) self.ref_grid = grid + self._ref_grid_params = ref_grid_params self.ref_grid.requires_grad = False return self.ref_grid diff --git a/tests/networks/blocks/warp/test_warp.py b/tests/networks/blocks/warp/test_warp.py index 1f23664234..ad30c76669 100644 --- a/tests/networks/blocks/warp/test_warp.py +++ b/tests/networks/blocks/warp/test_warp.py @@ -154,6 +154,41 @@ def test_jitter(self): self.assertTrue(torch.equal(same, repeat)) self.assertFalse(torch.equal(same, other)) + def test_reference_grid_cache(self): + """Verify reference-grid cache hits and invalidation across every key dimension.""" + warp_layer = Warp() + ddf = torch.zeros(1, 2, 4, 5) + + regular = warp_layer.get_reference_grid(ddf, jitter=False, seed=0) + self.assertIs(regular, warp_layer.get_reference_grid(ddf, jitter=False, seed=7)) + + float64 = warp_layer.get_reference_grid(ddf.to(torch.float64)) + self.assertIsNot(float64, regular) + self.assertEqual(float64.dtype, torch.float64) + + jitter_7 = warp_layer.get_reference_grid(ddf.to(torch.float64), jitter=True, seed=7) + self.assertIsNot(jitter_7, float64) + self.assertIs(jitter_7, warp_layer.get_reference_grid(ddf.to(torch.float64), jitter=True, seed=7)) + + jitter_8 = warp_layer.get_reference_grid(ddf.to(torch.float64), jitter=True, seed=8) + self.assertIsNot(jitter_8, jitter_7) + self.assertTrue(torch.equal(jitter_8, Warp().get_reference_grid(ddf.to(torch.float64), jitter=True, seed=8))) + + regular = warp_layer.get_reference_grid(ddf.to(torch.float64)) + self.assertIsNot(regular, jitter_8) + + different_batch = warp_layer.get_reference_grid(torch.zeros(2, 2, 4, 5, dtype=torch.float64)) + self.assertIsNot(different_batch, regular) + + different_shape = warp_layer.get_reference_grid(torch.zeros(2, 2, 5, 4, dtype=torch.float64)) + self.assertIsNot(different_shape, different_batch) + + meta_ddf = torch.zeros(2, 2, 5, 4, dtype=torch.float64, device="meta") + meta = warp_layer.get_reference_grid(meta_ddf) + self.assertIsNot(meta, different_shape) + self.assertIs(meta, warp_layer.get_reference_grid(meta_ddf)) + self.assertEqual(meta.device.type, "meta") + @mock.patch("monai.networks.blocks.warp.USE_COMPILED", False) def test_singleton_spatial_dim(self): """