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)
| 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 |
nothing calls this directly
no outgoing calls
no test coverage detected