MCPcopy Create free account
hub / github.com/Project-MONAI/MONAI / get_reference_grid

Method get_reference_grid

monai/networks/blocks/warp.py:114–131  ·  view source on GitHub ↗
(self, ddf: torch.Tensor, jitter: bool = False, seed: int = 0)

Source from the content-addressed store, hash-verified

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 """

Callers 1

forwardMethod · 0.95

Calls 1

meshgrid_ijFunction · 0.90

Tested by

no test coverage detected