Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 10 additions & 9 deletions fvdb_reality_capture/functional/_autograd.py
Original file line number Diff line number Diff line change
Expand Up @@ -412,8 +412,9 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor
class _RasterizeScreenSpaceGaussiansSparseFn(torch.autograd.Function):
"""Alpha-blending at an arbitrary set of pixels with gradients to the 2D quantities.

The forward pass takes and returns flat per-pixel tensors; the jagged structure of the pixel
selection is saved so the backward pass can rebuild the JaggedTensors the kernel expects.
The forward pass takes and returns flat per-pixel tensors. The pixel selection is kept on ``ctx``
so the backward pass can wrap the flat gradients in the JaggedTensors the kernel expects.
Rebuilding it from saved structure tensors would cost a device-to-host sync per backward step.
"""

@staticmethod
Expand Down Expand Up @@ -468,9 +469,6 @@ def forward(
opacities,
tile_offsets,
tile_gaussian_ids,
pixels_to_render.jdata,
pixels_to_render.joffsets,
pixels_to_render.jlidx,
alphas_jt.jdata,
last_ids_jt.jdata,
active_tiles,
Expand All @@ -481,6 +479,9 @@ def forward(
_save_optional(ctx, to_save, backgrounds, masks)
ctx.save_for_backward(*to_save)

# Integer pixel coordinates need no gradient and are not modified in place, so holding the
# JaggedTensor directly is safe
ctx.pixels_to_render = pixels_to_render
ctx.image_width = image_width
ctx.image_height = image_height
ctx.image_origin_w = image_origin_w
Expand All @@ -500,11 +501,11 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor

saved = ctx.saved_tensors
means2d, conics, features, opacities, tile_offsets, tile_gaussian_ids = saved[:6]
pixels_jdata, joffsets, jlidx, alphas_jdata, last_ids_jdata = saved[6:11]
active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map = saved[11:15]
backgrounds, masks = _load_optional(ctx, saved, 15)
alphas_jdata, last_ids_jdata = saved[6:8]
active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map = saved[8:12]
backgrounds, masks = _load_optional(ctx, saved, 12)

pixels_jt = JaggedTensor.from_data_offsets_and_list_ids(pixels_jdata, joffsets, jlidx)
pixels_jt: JaggedTensor = ctx.pixels_to_render
assert d_rendered is not None
assert d_alphas is not None
_, d_means2d, d_conics, d_features, d_opacities = F.rasterize_screen_space_gaussians_sparse_bwd(
Expand Down
Loading