Skip to content

Commit fd12fdc

Browse files
aymuos15ericspod
andauthored
Fix non-functional jitter in Warp.get_reference_grid (#8953)
### Description Little messy description as trying to fix multiple things but the jist is: `Warp.get_reference_grid` never applied the `jitter` it advertises and crashed whenever `jitter=True`. The grid is built from `torch.arange` (integer dtype) and `self.ref_grid` was assigned `grid.to(ddf)` before the jitter block, so `grid += torch.rand_like(grid)` mutated a local that was never returned, and `torch.rand_like` raises `NotImplementedError` on an integer tensor anyway. Separately, `fork_rng(enabled=seed)` disabled RNG forking when `seed` took its default of `0`, leaking the seeded state into the global RNG. The grid is now cast to `ddf` before jittering, the jittered tensor is assigned to `self.ref_grid`, and `fork_rng()` isolates the seeded draw. The non-jitter path is unchanged. A regression test covers the float/non-integer jittered grid, the integer un-jittered grid, and per-seed reproducibility; it fails before this change with `NotImplementedError`. ### Types of changes - [x] Non-breaking change (fix or new feature that would not break existing functionality). - [x] New tests added to cover the changes. --------- Signed-off-by: Soumya Snigdha Kundu <soumya_snigdha.kundu@kcl.ac.uk> Signed-off-by: Soumya Snigdha Kundu <soumyawork15@gmail.com> Signed-off-by: Eric Kerfoot <17726042+ericspod@users.noreply.github.com> Co-authored-by: Eric Kerfoot <17726042+ericspod@users.noreply.github.com>
1 parent 3bd4c4f commit fd12fdc

2 files changed

Lines changed: 18 additions & 2 deletions

File tree

monai/networks/blocks/warp.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -121,12 +121,13 @@ def get_reference_grid(self, ddf: torch.Tensor, jitter: bool = False, seed: int
121121
mesh_points = [torch.arange(0, dim) for dim in ddf.shape[2:]]
122122
grid = torch.stack(meshgrid_ij(*mesh_points), dim=0) # (spatial_dims, ...)
123123
grid = torch.stack([grid] * ddf.shape[0], dim=0) # (batch, spatial_dims, ...)
124-
self.ref_grid = grid.to(ddf)
124+
grid = grid.to(ddf)
125125
if jitter:
126126
# Define reference grid on non-integer values
127-
with torch.random.fork_rng(enabled=seed):
127+
with torch.random.fork_rng():
128128
torch.random.manual_seed(seed)
129129
grid += torch.rand_like(grid)
130+
self.ref_grid = grid
130131
self.ref_grid.requires_grad = False
131132
return self.ref_grid
132133

tests/networks/blocks/warp/test_warp.py

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -139,6 +139,21 @@ def test_ill_shape(self):
139139
with self.assertRaisesRegex(ValueError, ""):
140140
warp_layer(image=torch.arange(4).reshape((1, 1, 2, 2)).to(dtype=torch.float), ddf=torch.zeros(1, 2, 3, 3))
141141

142+
def test_jitter(self):
143+
ddf = torch.zeros(1, 2, 4, 5)
144+
grid = Warp(jitter=True).get_reference_grid(ddf, jitter=True, seed=0)
145+
self.assertTrue(grid.is_floating_point())
146+
self.assertFalse(torch.equal(grid, grid.round()))
147+
148+
grid = Warp().get_reference_grid(ddf, jitter=False)
149+
self.assertTrue(torch.equal(grid, grid.round()))
150+
151+
same = Warp().get_reference_grid(ddf, jitter=True, seed=7)
152+
repeat = Warp().get_reference_grid(ddf, jitter=True, seed=7)
153+
other = Warp().get_reference_grid(ddf, jitter=True, seed=8)
154+
self.assertTrue(torch.equal(same, repeat))
155+
self.assertFalse(torch.equal(same, other))
156+
142157
@mock.patch("monai.networks.blocks.warp.USE_COMPILED", False)
143158
def test_singleton_spatial_dim(self):
144159
"""

0 commit comments

Comments
 (0)