Skip to content
Merged
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
129 changes: 95 additions & 34 deletions fvdb_reality_capture/radiance_fields/gaussian_splatting.py
Original file line number Diff line number Diff line change
Expand Up @@ -1667,6 +1667,8 @@ def _intersect_tiles(
means2d: torch.Tensor,
radii: torch.Tensor,
depths: torch.Tensor,
conics: torch.Tensor,
opacities: torch.Tensor,
C: int,
tile_size: int,
W: int,
Expand All @@ -1686,6 +1688,8 @@ def _intersect_tiles(
tile_size,
num_tiles_h,
num_tiles_w,
conics=conics,
opacities=opacities,
)
return tile_offsets, tile_gaussian_ids, num_tiles_h, num_tiles_w

Expand All @@ -1695,6 +1699,8 @@ def _intersect_tiles_sparse(
means2d: torch.Tensor,
radii: torch.Tensor,
depths: torch.Tensor,
conics: torch.Tensor,
opacities: torch.Tensor,
C: int,
tile_size: int,
W: int,
Expand All @@ -1706,13 +1712,17 @@ def _intersect_tiles_sparse(
"""
num_tiles_h = math.ceil(H / tile_size)
num_tiles_w = math.ceil(W / tile_size)
active_tiles, active_tile_mask, tile_pixel_mask, tile_pixel_cumsum, pixel_map = (
_C.build_sparse_gaussian_tile_layout(
tile_size,
num_tiles_w,
num_tiles_h,
pixels_jt._impl,
)
(
active_tiles,
active_tile_mask,
tile_pixel_mask,
tile_pixel_cumsum,
pixel_map,
) = _C.build_sparse_gaussian_tile_layout(
tile_size,
num_tiles_w,
num_tiles_h,
pixels_jt._impl,
)
tile_offsets, tile_gaussian_ids = _C.intersect_gaussian_tiles_sparse(
means2d,
Expand All @@ -1724,6 +1734,8 @@ def _intersect_tiles_sparse(
tile_size,
num_tiles_h,
num_tiles_w,
conics=conics,
opacities=opacities,
)
return tile_offsets, tile_gaussian_ids, active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map

Expand Down Expand Up @@ -1961,9 +1973,14 @@ def _sparse_render_impl(
opacities = self._make_opacities(C, compensations, antialias)
features = self._make_render_features(w2c, radii, depths, sh_degree_to_use, include_colors, include_depth)

tile_offsets, tile_gaussian_ids, active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map = (
self._intersect_tiles_sparse(render_pixels, means2d, radii, depths, C, tile_size, W, H)
)
(
tile_offsets,
tile_gaussian_ids,
active_tiles,
tile_pixel_mask,
tile_pixel_cumsum,
pixel_map,
) = self._intersect_tiles_sparse(render_pixels, means2d, radii, depths, conics, opacities, C, tile_size, W, H)

rendered_jdata, alphas_jdata = self._rasterize_screen_space_sparse(
render_pixels,
Expand Down Expand Up @@ -2503,6 +2520,8 @@ def render_from_projected_gaussians(
tile_size,
num_tiles_h,
num_tiles_w,
conics=pg.inv_covar_2d,
opacities=pg.opacities,
)
features, alphas = cast(
tuple[torch.Tensor, torch.Tensor],
Expand All @@ -2528,6 +2547,8 @@ def render_from_projected_gaussians(
pg.means2d,
pg.radii,
pg.depths,
pg.inv_covar_2d,
pg.opacities,
C,
tile_size,
W,
Expand Down Expand Up @@ -2655,6 +2676,8 @@ def render_depths(
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
Expand Down Expand Up @@ -2903,6 +2926,8 @@ def render_images(
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
Expand Down Expand Up @@ -3041,6 +3066,8 @@ def render_images_from_world(
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
Expand Down Expand Up @@ -3115,6 +3142,8 @@ def render_depths_from_world(
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
Expand Down Expand Up @@ -3509,6 +3538,8 @@ def render_images_and_depths(
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
Expand Down Expand Up @@ -3585,6 +3616,8 @@ def render_images_and_depths_from_world(
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
Expand Down Expand Up @@ -3711,6 +3744,8 @@ def render_num_contributing_gaussians(
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
Expand Down Expand Up @@ -3857,17 +3892,24 @@ def sparse_render_num_contributing_gaussians(
)
C = world_to_camera_matrices.size(0)
opacities = self._make_opacities(C, compensations, antialias)
tile_offsets, tile_gaussian_ids, active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map = (
self._intersect_tiles_sparse(
unique_pixels_jt,
means2d,
radii,
depths,
C,
tile_size,
image_width,
image_height,
)
(
tile_offsets,
tile_gaussian_ids,
active_tiles,
tile_pixel_mask,
tile_pixel_cumsum,
pixel_map,
) = self._intersect_tiles_sparse(
unique_pixels_jt,
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
image_height,
)
result_ncg, result_alphas = _C.sparse_rasterize_num_contributing_gaussians(
means2d,
Expand Down Expand Up @@ -3983,6 +4025,8 @@ def render_contributing_gaussian_ids(
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
Expand Down Expand Up @@ -4144,17 +4188,24 @@ def sparse_render_contributing_gaussian_ids(
)
C = world_to_camera_matrices.size(0)
opacities = self._make_opacities(C, compensations, antialias)
tile_offsets, tile_gaussian_ids, active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map = (
self._intersect_tiles_sparse(
unique_pixels_jt,
means2d,
radii,
depths,
C,
tile_size,
image_width,
image_height,
)
(
tile_offsets,
tile_gaussian_ids,
active_tiles,
tile_pixel_mask,
tile_pixel_cumsum,
pixel_map,
) = self._intersect_tiles_sparse(
unique_pixels_jt,
means2d,
radii,
depths,
conics,
opacities,
C,
tile_size,
image_width,
image_height,
)
ncg_jt = None
if top_k_contributors <= 0:
Expand Down Expand Up @@ -4669,6 +4720,7 @@ def gaussian_render_jagged(
opacities_batched = opacities.jdata[gaussian_ids]
if antialias:
opacities_batched = opacities_batched * compensations
opacities_batched = opacities_batched.contiguous()

debug_info: dict[str, torch.Tensor] = {}
if return_debug_info:
Expand Down Expand Up @@ -4723,7 +4775,16 @@ def gaussian_render_jagged(
num_tiles_h = math.ceil(image_height / tile_size)
num_tiles_w = math.ceil(image_width / tile_size)
tile_offsets, tile_gaussian_ids_t = _C.intersect_gaussian_tiles(
means2d, radii, depths, ccz, tile_size, num_tiles_h, num_tiles_w, camera_ids
means2d,
radii,
depths,
ccz,
tile_size,
num_tiles_h,
num_tiles_w,
camera_ids=camera_ids,
conics=conics,
opacities=opacities_batched,
)
if return_debug_info:
debug_info["tile_offsets"] = tile_offsets
Expand All @@ -4734,7 +4795,7 @@ def gaussian_render_jagged(
means2d,
conics,
render_quantities,
opacities_batched.contiguous(),
opacities_batched,
image_width,
image_height,
0, # image_origin_w
Expand Down
Loading