Skip to content

Commit 1d7d772

Browse files
committed
Fix Warp reference grid caching
1 parent 605611b commit 1d7d772

2 files changed

Lines changed: 42 additions & 2 deletions

File tree

monai/networks/blocks/warp.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -109,13 +109,17 @@ def __init__(self, mode=GridSampleMode.BILINEAR.value, padding_mode=GridSamplePa
109109
self._padding_mode = self._padding_mode_native
110110

111111
self.ref_grid = None
112+
self._ref_grid_params: tuple[bool, int | None] | None = None
112113
self.jitter = jitter
113114

114115
def get_reference_grid(self, ddf: torch.Tensor, jitter: bool = False, seed: int = 0) -> torch.Tensor:
116+
ref_grid_params = (jitter, seed if jitter else None)
115117
if (
116118
self.ref_grid is not None
117-
and self.ref_grid.shape[0] == ddf.shape[0]
118-
and self.ref_grid.shape[1:] == ddf.shape[2:]
119+
and self.ref_grid.shape == ddf.shape
120+
and self.ref_grid.device == ddf.device
121+
and self.ref_grid.dtype == ddf.dtype
122+
and self._ref_grid_params == ref_grid_params
119123
):
120124
return self.ref_grid # type: ignore
121125
mesh_points = [torch.arange(0, dim) for dim in ddf.shape[2:]]
@@ -128,6 +132,7 @@ def get_reference_grid(self, ddf: torch.Tensor, jitter: bool = False, seed: int
128132
torch.random.manual_seed(seed)
129133
grid += torch.rand_like(grid)
130134
self.ref_grid = grid
135+
self._ref_grid_params = ref_grid_params
131136
self.ref_grid.requires_grad = False
132137
return self.ref_grid
133138

tests/networks/blocks/warp/test_warp.py

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -154,6 +154,41 @@ def test_jitter(self):
154154
self.assertTrue(torch.equal(same, repeat))
155155
self.assertFalse(torch.equal(same, other))
156156

157+
def test_reference_grid_cache(self):
158+
"""Verify reference-grid cache hits and invalidation across every key dimension."""
159+
warp_layer = Warp()
160+
ddf = torch.zeros(1, 2, 4, 5)
161+
162+
regular = warp_layer.get_reference_grid(ddf, jitter=False, seed=0)
163+
self.assertIs(regular, warp_layer.get_reference_grid(ddf, jitter=False, seed=7))
164+
165+
float64 = warp_layer.get_reference_grid(ddf.to(torch.float64))
166+
self.assertIsNot(float64, regular)
167+
self.assertEqual(float64.dtype, torch.float64)
168+
169+
jitter_7 = warp_layer.get_reference_grid(ddf.to(torch.float64), jitter=True, seed=7)
170+
self.assertIsNot(jitter_7, float64)
171+
self.assertIs(jitter_7, warp_layer.get_reference_grid(ddf.to(torch.float64), jitter=True, seed=7))
172+
173+
jitter_8 = warp_layer.get_reference_grid(ddf.to(torch.float64), jitter=True, seed=8)
174+
self.assertIsNot(jitter_8, jitter_7)
175+
self.assertTrue(torch.equal(jitter_8, Warp().get_reference_grid(ddf.to(torch.float64), jitter=True, seed=8)))
176+
177+
regular = warp_layer.get_reference_grid(ddf.to(torch.float64))
178+
self.assertIsNot(regular, jitter_8)
179+
180+
different_batch = warp_layer.get_reference_grid(torch.zeros(2, 2, 4, 5, dtype=torch.float64))
181+
self.assertIsNot(different_batch, regular)
182+
183+
different_shape = warp_layer.get_reference_grid(torch.zeros(2, 2, 5, 4, dtype=torch.float64))
184+
self.assertIsNot(different_shape, different_batch)
185+
186+
meta_ddf = torch.zeros(2, 2, 5, 4, dtype=torch.float64, device="meta")
187+
meta = warp_layer.get_reference_grid(meta_ddf)
188+
self.assertIsNot(meta, different_shape)
189+
self.assertIs(meta, warp_layer.get_reference_grid(meta_ddf))
190+
self.assertEqual(meta.device.type, "meta")
191+
157192
@mock.patch("monai.networks.blocks.warp.USE_COMPILED", False)
158193
def test_singleton_spatial_dim(self):
159194
"""

0 commit comments

Comments
 (0)