| 112 | self.jitter = jitter |
| 113 | |
| 114 | def get_reference_grid(self, ddf: torch.Tensor, jitter: bool = False, seed: int = 0) -> torch.Tensor: |
| 115 | if ( |
| 116 | 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 | ): |
| 120 | return self.ref_grid # type: ignore |
| 121 | mesh_points = [torch.arange(0, dim) for dim in ddf.shape[2:]] |
| 122 | grid = torch.stack(meshgrid_ij(*mesh_points), dim=0) # (spatial_dims, ...) |
| 123 | grid = torch.stack([grid] * ddf.shape[0], dim=0) # (batch, spatial_dims, ...) |
| 124 | self.ref_grid = grid.to(ddf) |
| 125 | if jitter: |
| 126 | # Define reference grid on non-integer values |
| 127 | with torch.random.fork_rng(enabled=seed): |
| 128 | torch.random.manual_seed(seed) |
| 129 | grid += torch.rand_like(grid) |
| 130 | self.ref_grid.requires_grad = False |
| 131 | return self.ref_grid |
| 132 | |
| 133 | def forward(self, image: torch.Tensor, ddf: torch.Tensor): |
| 134 | """ |