MCPcopy Create free account
hub / github.com/JasonLSC/GSCodec_Studio / forward

Method forward

examples/lib_bilagrid.py:310–365  ·  view source on GitHub ↗

Bilateral grid slicing. Supports 2-D, 3-D, 4-D, and 5-D input. For the 2-D, 3-D, and 4-D cases, please refer to `slice`. For the 5-D cases, `idx` will be unused and the first dimension of `xy` should be equal to the number of bilateral grids. Then this function becomes PyTorc

(self, grid_xy, rgb, idx=None)

Source from the content-addressed store, hash-verified

308 return total_variation_loss(self.grids)
309
310 def forward(self, grid_xy, rgb, idx=None):
311 """Bilateral grid slicing. Supports 2-D, 3-D, 4-D, and 5-D input.
312 For the 2-D, 3-D, and 4-D cases, please refer to `slice`.
313 For the 5-D cases, `idx` will be unused and the first dimension of `xy` should be
314 equal to the number of bilateral grids. Then this function becomes PyTorch's
315 [`F.grid_sample`](https://pytorch.org/docs/stable/generated/torch.nn.functional.grid_sample.html).
316
317 Args:
318 grid_xy (torch.Tensor): The x-y coordinates in the range of $[0,1]$.
319 rgb (torch.Tensor): The RGB values in the range of $[0,1]$.
320 idx (torch.Tensor): The bilateral grid indices.
321
322 Returns:
323 Sliced affine matrices of shape $(..., 3, 4)$.
324 """
325
326 grids = self.grids
327 input_ndims = len(grid_xy.shape)
328 assert len(rgb.shape) == input_ndims
329
330 if input_ndims > 1 and input_ndims < 5:
331 # Convert input into 5D
332 for i in range(5 - input_ndims):
333 grid_xy = grid_xy.unsqueeze(1)
334 rgb = rgb.unsqueeze(1)
335 assert idx is not None
336 elif input_ndims != 5:
337 raise ValueError(
338 "Bilateral grid slicing only takes either 2D, 3D, 4D and 5D inputs"
339 )
340
341 grids = self.grids
342 if idx is not None:
343 grids = grids[idx]
344 assert grids.shape[0] == grid_xy.shape[0]
345
346 # Generate slicing coordinates.
347 grid_xy = (grid_xy - 0.5) * 2 # Rescale to [-1, 1].
348 grid_z = self.rgb2gray(rgb)
349
350 # print(grid_xy.shape, grid_z.shape)
351 # exit()
352 grid_xyz = torch.cat([grid_xy, grid_z], dim=-1) # (N, m, h, w, 3)
353
354 affine_mats = F.grid_sample(
355 grids, grid_xyz, mode="bilinear", align_corners=True, padding_mode="border"
356 ) # (N, 12, m, h, w)
357 affine_mats = affine_mats.permute(0, 2, 3, 4, 1) # (N, m, h, w, 12)
358 affine_mats = affine_mats.reshape(
359 *affine_mats.shape[:-1], 3, 4
360 ) # (N, m, h, w, 3, 4)
361
362 for _ in range(5 - input_ndims):
363 affine_mats = affine_mats.squeeze(1)
364
365 return affine_mats
366
367

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected