Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 17 additions & 2 deletions monai/networks/blocks/warp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:]]
Expand All @@ -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

Expand Down
35 changes: 35 additions & 0 deletions tests/networks/blocks/warp/test_warp.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
"""
Expand Down
Loading