diff --git a/docs/api/enums.rst b/docs/api/enums.rst index f7927361..fce7da4c 100644 --- a/docs/api/enums.rst +++ b/docs/api/enums.rst @@ -1,15 +1,18 @@ Enums ===== -The camera and Gaussian-splatting enums are provided by -``fvdb_reality_capture``. Their member values remain compatible with the -underlying compiled fVDB kernels. +The camera enums that fvdb kernels accept are owned by ``fvdb`` and re-exported by +``fvdb_reality_capture`` as the same objects, so values pass between the two packages without +conversion: -.. autoclass:: fvdb_reality_capture.RollingShutterType - :members: +- :class:`fvdb.CameraModel` (also available as ``fvdb_reality_capture.CameraModel``) +- :class:`fvdb.RollingShutterType` (also available as ``fvdb_reality_capture.RollingShutterType``) -.. autoclass:: fvdb_reality_capture.CameraModel - :members: +:class:`ProjectionMethod` and :class:`GaussianRenderMode` select stages of the composable rendering +pipeline in :mod:`fvdb_reality_capture.functional` and are defined here. .. autoclass:: fvdb_reality_capture.ProjectionMethod :members: + +.. autoclass:: fvdb_reality_capture.GaussianRenderMode + :members: diff --git a/docs/api/functional.rst b/docs/api/functional.rst new file mode 100644 index 00000000..16dc8344 --- /dev/null +++ b/docs/api/functional.rst @@ -0,0 +1,142 @@ +Functional Gaussian Splatting +============================= + +.. module:: fvdb_reality_capture.functional + +:mod:`fvdb_reality_capture.functional` exposes Gaussian splat rendering as four composable stages +that pass small frozen dataclasses between them. :class:`~fvdb_reality_capture.GaussianSplat3d` +is a thin wrapper that composes these stages; use them directly when you need to insert your own +logic between projection and rasterization, reuse a projection for several renders, or build a +training loop over plain tensors. + +The kernels themselves live in :mod:`fvdb.functional` as flat, non-differentiable forward and +backward functions. The stages here attach autograd to them, so gradients flow through every stage +except tile intersection. + +.. code-block:: python + + import torch + import fvdb_reality_capture.functional as F + from fvdb_reality_capture import CameraModel, GaussianRenderMode + + # Plain tensors: means [N, 3], quats [N, 4], log_scales [N, 3], logit_opacities [N], + # sh0 [N, 1, 3], shN [N, K - 1, 3], world_to_cam [C, 4, 4], K [C, 3, 3] + + # Stage 1: project the 3D Gaussians into every camera + projected = F.project_gaussians( + means, quats, log_scales, world_to_cam, K, image_width=640, image_height=480 + ) + + # Stage 2: view-dependent features from spherical harmonics + features = F.evaluate_gaussian_sh( + means, sh0, shN, world_to_cam, projected, render_mode=GaussianRenderMode.FEATURES + ) + + # Per-camera opacities, computed once and shared by stages 3 and 4 + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + + # Stage 3: bin Gaussians into image tiles (opacities enable tighter culling) + tiles = F.intersect_gaussian_tiles(projected, opacities) + + # Stage 4: alpha-blend into images + images, alphas = F.rasterize_screen_space_gaussians(projected, features, opacities, tiles) + + loss = torch.nn.functional.l1_loss(images, target_images) + loss.backward() # gradients reach means, quats, log_scales, logit_opacities, sh0 and shN + +The sparse path renders an arbitrary set of pixels per camera. Swap stages 3 and 4 for +:func:`intersect_gaussian_tiles_sparse` and :func:`rasterize_screen_space_gaussians_sparse`; results +come back as :class:`~fvdb.JaggedTensor` in the order of the requested pixels, duplicates included. +The world-space path, :func:`rasterize_world_space_gaussians`, evaluates the 3D Gaussians along +per-pixel rays and is the training path for the unscented projection, whose kernel has no backward +pass. + + +Types +----- + +.. autoclass:: ProjectedGaussians + :members: + +.. autoclass:: GaussianTileIntersection + :members: + +.. autoclass:: SparseGaussianTileIntersection + :members: + + +Stage 1: Projection +------------------- + +.. autofunction:: project_gaussians + +.. autofunction:: resolve_projection_method + +.. autofunction:: requires_distortion_coeffs +.. autofunction:: check_distortion_coeffs + + +Stage 2: Features +----------------- + +.. autofunction:: evaluate_gaussian_sh + +.. autofunction:: sh_degree_from_coefficients + + +Stage 3: Tile Intersection +-------------------------- + +.. autofunction:: intersect_gaussian_tiles + +.. autofunction:: intersect_gaussian_tiles_sparse + +.. autofunction:: deduplicate_pixels + +.. autofunction:: check_tiles_match + +.. py:function:: as_pixel_jagged(value) + + ``fvdb.functional.as_pixel_jagged``, re-exported so the sparse pipeline here applies the same + pixel-selection checks as fvdb's own sparse kernels. Normalizes a ``[C, P, 2]`` tensor or a + :class:`~fvdb.JaggedTensor` of ``(row, col)`` integer pixels to a JaggedTensor with one list per + camera, raising ``TypeError`` for non-integer coordinates and ``ValueError`` for malformed shapes. + + +Stage 4: Rasterization +---------------------- + +.. autofunction:: rasterize_screen_space_gaussians + +.. autofunction:: rasterize_world_space_gaussians + +.. autofunction:: rasterize_screen_space_gaussians_sparse + +.. autofunction:: compute_gaussian_opacities + +.. py:data:: Crop + + ``tuple[int, int, int, int]``: a crop window as ``(origin_w, origin_h, width, height)`` in pixels. + +.. autofunction:: validate_crop + +.. autofunction:: apply_crop +.. autofunction:: apply_pixel_mask + +.. autofunction:: pad_crop + +.. autofunction:: pixel_mask_to_tile_mask + + +Analysis +-------- + +These do not build an autograd graph. + +.. autofunction:: rasterize_num_contributing_gaussians + +.. autofunction:: rasterize_contributing_gaussian_ids + +.. autofunction:: rasterize_num_contributing_gaussians_sparse + +.. autofunction:: rasterize_contributing_gaussian_ids_sparse diff --git a/docs/api/gaussian_splatting.rst b/docs/api/gaussian_splatting.rst index 8899671b..69f56338 100644 --- a/docs/api/gaussian_splatting.rst +++ b/docs/api/gaussian_splatting.rst @@ -2,8 +2,9 @@ Gaussian Splatting ================== The high-level Gaussian splatting API is provided by -``fvdb_reality_capture``. The underlying rendering kernels and supporting -tensor types remain in ``fvdb``. +``fvdb_reality_capture``. :class:`GaussianSplat3d` composes the stages of +:mod:`fvdb_reality_capture.functional`; the underlying rendering kernels and +supporting tensor types remain in ``fvdb``. .. autoclass:: fvdb_reality_capture.ProjectedGaussianSplats :members: diff --git a/docs/conf.py b/docs/conf.py index 848f626f..ab14f998 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -37,7 +37,16 @@ # Add any Sphinx extension module names here, as strings. They can be # extensions coming with Sphinx (named 'sphinx.ext.*') or your custom # ones. -extensions = ["sphinx.ext.autodoc", "sphinx.ext.viewcode", "sphinx.ext.napoleon", "myst_parser"] +extensions = [ + "sphinx.ext.autodoc", + "sphinx.ext.intersphinx", + "sphinx.ext.viewcode", + "sphinx.ext.napoleon", + "myst_parser", +] + +# fvdb is mocked during the docs build; resolve references to it against its published docs. +intersphinx_mapping = {"fvdb": ("https://fvdb-core.readthedocs.io/latest/", None)} myst_enable_extensions = [ "amsmath", diff --git a/docs/index.rst b/docs/index.rst index 63011503..b5426610 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -105,6 +105,7 @@ A common reality capture pipeline typically resembles the figure below: api/enums api/gaussian_splatting + api/functional api/radiance_fields api/sfm_scene api/tools diff --git a/fvdb_reality_capture/__init__.py b/fvdb_reality_capture/__init__.py index 913b1779..c759b2da 100644 --- a/fvdb_reality_capture/__init__.py +++ b/fvdb_reality_capture/__init__.py @@ -14,7 +14,8 @@ tools, transforms, ) -from .enums import CameraModel, ProjectionMethod, RollingShutterType +from . import functional +from .enums import CameraModel, GaussianRenderMode, ProjectionMethod, RollingShutterType from .radiance_fields import ( GaussianSplat3d, ProjectedGaussianSplats, @@ -39,6 +40,8 @@ "RollingShutterType", "CameraModel", "ProjectionMethod", + "GaussianRenderMode", + "functional", "checkpoints", "dev", "foundation_models", diff --git a/fvdb_reality_capture/enums.py b/fvdb_reality_capture/enums.py index b02e459a..bf737bc8 100644 --- a/fvdb_reality_capture/enums.py +++ b/fvdb_reality_capture/enums.py @@ -1,99 +1,47 @@ # Copyright Contributors to the OpenVDB Project # SPDX-License-Identifier: Apache-2.0 # +"""Enums used by the Gaussian splatting API. -from enum import IntEnum - -__all__ = ["RollingShutterType", "CameraModel", "ProjectionMethod"] - - -class RollingShutterType(IntEnum): - """ - Rolling shutter policy for camera projection / ray generation. - - Rolling shutter models treat different image rows/columns as having different exposure times. - FVDB uses this to interpolate between per-camera start/end poses when generating rays. - """ - - NONE = 0 - """ - No rolling shutter: the start pose is used for all pixels. - """ - - VERTICAL = 1 - """ - Vertical rolling shutter: exposure time varies with image row (y). - """ - - HORIZONTAL = 2 - """ - Horizontal rolling shutter: exposure time varies with image column (x). - """ - - -class CameraModel(IntEnum): - """ - Camera model for projection / ray generation. - - Notes: +The camera enums that fvdb kernels accept are owned by :mod:`fvdb` and re-exported here unchanged, +so :class:`fvdb_reality_capture.CameraModel` is the same object as :class:`fvdb.CameraModel`. +:class:`ProjectionMethod` and :class:`GaussianRenderMode` select stages of the composable rendering +pipeline in :mod:`fvdb_reality_capture.functional` and are defined here. +""" - - ``PINHOLE`` and ``ORTHOGRAPHIC`` ignore distortion coefficients. - - ``OPENCV_*`` variants use pinhole intrinsics plus OpenCV-style distortion. When distortion - coefficients are provided, FVDB expects a packed layout: +from enum import IntEnum - ``[k1,k2,k3,k4,k5,k6,p1,p2,s1,s2,s3,s4]`` +from fvdb import CameraModel, RollingShutterType - Unused coefficients for a given model should be set to 0. - """ +__all__ = ["RollingShutterType", "CameraModel", "ProjectionMethod", "GaussianRenderMode"] - PINHOLE = 0 - """ - Ideal pinhole camera model (no distortion). - """ - OPENCV_RADTAN_5 = 1 - """ - OpenCV radial-tangential distortion with 5 parameters (k1,k2,p1,p2,k3). - """ - - OPENCV_RATIONAL_8 = 2 +class ProjectionMethod(IntEnum): """ - OpenCV rational radial-tangential distortion with 8 parameters (k1..k6,p1,p2). + Which fvdb projection kernel :func:`fvdb_reality_capture.functional.project_gaussians` calls. """ - OPENCV_RADTAN_THIN_PRISM_9 = 3 - """ - OpenCV radial-tangential + thin-prism distortion with 9 parameters (k1,k2,p1,p2,k3,s1..s4). - """ + AUTO = 0 + """Choose the default implementation for the selected camera model.""" - OPENCV_THIN_PRISM_12 = 4 - """ - OpenCV rational radial-tangential + thin-prism distortion with 12 parameters - (k1..k6,p1,p2,s1..s4). - """ + ANALYTIC = 1 + """Use the analytic (EWA) projection path.""" - ORTHOGRAPHIC = 5 - """ - Orthographic camera model (no distortion). - """ + UNSCENTED = 2 + """Use the unscented-transform projection path.""" -class ProjectionMethod(IntEnum): +class GaussianRenderMode(IntEnum): """ - Projection implementation selector for Gaussian splatting camera models. + Which per-Gaussian features :func:`fvdb_reality_capture.functional.evaluate_gaussian_sh` produces + for rasterization. """ - AUTO = 0 - """ - Choose the default implementation for the selected camera model. - """ + FEATURES = 0 + """Spherical-harmonics evaluated features only, ``[C, N, D]``.""" - ANALYTIC = 1 - """ - Use the analytic projection path. - """ + DEPTH = 1 + """View-space depth only, ``[C, N, 1]``. No spherical harmonics are evaluated.""" - UNSCENTED = 2 - """ - Use the unscented projection path. - """ + FEATURES_AND_DEPTH = 2 + """Spherical-harmonics features with depth appended as the last channel, ``[C, N, D + 1]``.""" diff --git a/fvdb_reality_capture/functional/__init__.py b/fvdb_reality_capture/functional/__init__.py new file mode 100644 index 00000000..bba42682 --- /dev/null +++ b/fvdb_reality_capture/functional/__init__.py @@ -0,0 +1,92 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +""" +``fvdb_reality_capture.functional`` -- composable, differentiable Gaussian splatting. + +Rendering is split into four stages that pass small frozen dataclasses between them, so a custom +pipeline can insert its own logic anywhere: + +1. :func:`project_gaussians` projects the 3D Gaussians into every camera. +2. :func:`evaluate_gaussian_sh` turns spherical-harmonics coefficients into per-camera features. +3. :func:`intersect_gaussian_tiles` or :func:`intersect_gaussian_tiles_sparse` bins Gaussians into tiles. +4. :func:`rasterize_screen_space_gaussians`, :func:`rasterize_world_space_gaussians` or + :func:`rasterize_screen_space_gaussians_sparse` alpha-blends the features. + +Stages 3 and 4 take per-camera opacities, ``[C, N]``. Compute them once per render with +:func:`compute_gaussian_opacities` and pass the same tensor to every stage. + +Every stage except tile intersection is differentiable. The kernels themselves live in +:mod:`fvdb.functional`; :class:`~fvdb_reality_capture.GaussianSplat3d` composes these stages. +""" + +from ._analysis import ( + rasterize_contributing_gaussian_ids, + rasterize_contributing_gaussian_ids_sparse, + rasterize_num_contributing_gaussians, + rasterize_num_contributing_gaussians_sparse, +) +from ._opacity import compute_gaussian_opacities +from ._projection import ( + check_distortion_coeffs, + project_gaussians, + requires_distortion_coeffs, + resolve_projection_method, +) +from ._rasterization import ( + Crop, + apply_crop, + apply_pixel_mask, + pad_crop, + pixel_mask_to_tile_mask, + rasterize_screen_space_gaussians, + rasterize_screen_space_gaussians_sparse, + rasterize_world_space_gaussians, + validate_crop, +) +from ._spherical_harmonics import evaluate_gaussian_sh, sh_degree_from_coefficients +from ._tile_intersection import ( + as_pixel_jagged, + check_tiles_match, + deduplicate_pixels, + intersect_gaussian_tiles, + intersect_gaussian_tiles_sparse, +) +from ._types import GaussianTileIntersection, ProjectedGaussians, SparseGaussianTileIntersection + +__all__ = [ + # Types + "ProjectedGaussians", + "GaussianTileIntersection", + "SparseGaussianTileIntersection", + # Stage 1: projection + "project_gaussians", + "resolve_projection_method", + "check_distortion_coeffs", + "requires_distortion_coeffs", + # Stage 2: features + "evaluate_gaussian_sh", + "sh_degree_from_coefficients", + # Stage 3: tile intersection + "intersect_gaussian_tiles", + "intersect_gaussian_tiles_sparse", + "deduplicate_pixels", + "as_pixel_jagged", + "check_tiles_match", + # Stage 4: rasterization + "rasterize_screen_space_gaussians", + "rasterize_world_space_gaussians", + "rasterize_screen_space_gaussians_sparse", + "compute_gaussian_opacities", + "pixel_mask_to_tile_mask", + "Crop", + "validate_crop", + "apply_crop", + "apply_pixel_mask", + "pad_crop", + # Analysis + "rasterize_num_contributing_gaussians", + "rasterize_num_contributing_gaussians_sparse", + "rasterize_contributing_gaussian_ids", + "rasterize_contributing_gaussian_ids_sparse", +] diff --git a/fvdb_reality_capture/functional/_analysis.py b/fvdb_reality_capture/functional/_analysis.py new file mode 100644 index 00000000..dcc7f75e --- /dev/null +++ b/fvdb_reality_capture/functional/_analysis.py @@ -0,0 +1,245 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +"""Non-differentiable analysis of which Gaussians contribute to which pixels.""" + +from __future__ import annotations + +import torch +from fvdb import JaggedTensor +from fvdb import functional as F + +from ._opacity import check_opacities +from ._tile_intersection import check_tiles_match +from ._types import GaussianTileIntersection, ProjectedGaussians, SparseGaussianTileIntersection + + +def rasterize_num_contributing_gaussians( + projected: ProjectedGaussians, + opacities: torch.Tensor, + tiles: GaussianTileIntersection, +) -> tuple[torch.Tensor, torch.Tensor]: + """Count the Gaussians that contribute non-negligible opacity to each pixel. + + Args: + projected (ProjectedGaussians): Output of :func:`project_gaussians`. + opacities (torch.Tensor): Per-camera opacities, ``[C, N]``, from :func:`compute_gaussian_opacities`. + tiles (GaussianTileIntersection): Output of :func:`intersect_gaussian_tiles` for ``projected``. + + Returns: + num_contributing (torch.Tensor): Contributor count per pixel, ``[C, H, W]``, ``int32``. + alphas (torch.Tensor): Accumulated alpha per pixel, ``[C, H, W]``. + """ + with torch.no_grad(): + return _count_dense(projected, check_opacities(opacities, projected), tiles) + + +def _count_dense( + projected: ProjectedGaussians, opacities: torch.Tensor, tiles: GaussianTileIntersection +) -> tuple[torch.Tensor, torch.Tensor]: + check_tiles_match(tiles, projected) + return F.rasterize_num_contributing_gaussians( + projected.means2d, + projected.conics, + opacities, + tiles.tile_offsets, + tiles.tile_gaussian_ids, + tiles.image_width, + tiles.image_height, + 0, + 0, + tiles.tile_size, + ) + + +def rasterize_contributing_gaussian_ids( + projected: ProjectedGaussians, + opacities: torch.Tensor, + tiles: GaussianTileIntersection, + top_k_contributors: int = 0, + num_contributing: torch.Tensor | None = None, +) -> tuple[JaggedTensor, JaggedTensor]: + """List the Gaussians contributing to each pixel, front to back, with their blend weights. + + With ``top_k_contributors > 0`` at most that many of the most visible contributors are kept per + pixel. Otherwise every contributor is listed, which needs the per-pixel counts; they are computed + here unless ``num_contributing`` supplies them. + + Args: + projected (ProjectedGaussians): Output of :func:`project_gaussians`. + opacities (torch.Tensor): Per-camera opacities, ``[C, N]``, from :func:`compute_gaussian_opacities`. + tiles (GaussianTileIntersection): Output of :func:`intersect_gaussian_tiles` for ``projected``. + top_k_contributors (int): Contributors to keep per pixel, or ``0`` for all of them. + num_contributing (torch.Tensor | None): Counts from :func:`rasterize_num_contributing_gaussians`, + ``[C, H, W]``. Used only when ``top_k_contributors <= 0``. + + Returns: + gaussian_ids (JaggedTensor): Contributor indices, nested as cameras, then pixels, then contributors. + weights (JaggedTensor): Blend weight of each listed contributor, same structure. + """ + with torch.no_grad(): + opacities = check_opacities(opacities, projected) + check_tiles_match(tiles, projected) + if top_k_contributors <= 0 and num_contributing is None: + num_contributing, _ = _count_dense(projected, opacities, tiles) + return F.rasterize_contributing_gaussian_ids( + projected.means2d, + projected.conics, + opacities, + tiles.tile_offsets, + tiles.tile_gaussian_ids, + tiles.image_width, + tiles.image_height, + 0, + 0, + tiles.tile_size, + top_k_contributors, + num_contributing if top_k_contributors <= 0 else None, + ) + + +def rasterize_num_contributing_gaussians_sparse( + projected: ProjectedGaussians, + opacities: torch.Tensor, + sparse_tiles: SparseGaussianTileIntersection, +) -> tuple[JaggedTensor, JaggedTensor]: + """Count contributing Gaussians at the requested pixels only. + + Args: + projected (ProjectedGaussians): Output of :func:`project_gaussians`. + opacities (torch.Tensor): Per-camera opacities, ``[C, N]``, from :func:`compute_gaussian_opacities`. + sparse_tiles (SparseGaussianTileIntersection): Output of :func:`intersect_gaussian_tiles_sparse`. + + Returns: + num_contributing (JaggedTensor): Contributor count per requested pixel, ``int32``, one list per camera. + alphas (JaggedTensor): Accumulated alpha per requested pixel, one list per camera. + """ + with torch.no_grad(): + counts, alphas = _count_sparse_unique(projected, check_opacities(opacities, projected), sparse_tiles) + requested = sparse_tiles.pixels_to_render + return ( + requested.jagged_like(sparse_tiles.expand_to_requested(counts.jdata)), + requested.jagged_like(sparse_tiles.expand_to_requested(alphas.jdata)), + ) + + +def _count_sparse_unique( + projected: ProjectedGaussians, opacities: torch.Tensor, sparse_tiles: SparseGaussianTileIntersection +) -> tuple[JaggedTensor, JaggedTensor]: + """Contributor counts and alphas over the unique pixels, as the kernel produces them.""" + check_tiles_match(sparse_tiles, projected) + return F.rasterize_num_contributing_gaussians_sparse( + projected.means2d, + projected.conics, + opacities, + sparse_tiles.tile_offsets, + sparse_tiles.tile_gaussian_ids, + sparse_tiles.unique_pixels, + sparse_tiles.active_tiles, + sparse_tiles.tile_pixel_mask, + sparse_tiles.tile_pixel_cumsum, + sparse_tiles.pixel_map, + sparse_tiles.image_width, + sparse_tiles.image_height, + 0, + 0, + sparse_tiles.tile_size, + ) + + +def _expand_contributions( + sparse_tiles: SparseGaussianTileIntersection, ids: JaggedTensor, weights: JaggedTensor +) -> tuple[JaggedTensor, JaggedTensor]: + """Expand per-unique-pixel contributor lists to the requested pixels, repeating duplicates. + + ``ids`` and ``weights`` list the same contributors, so one gather plan (built with a single + device-to-host sync for the total count) serves both. The results nest cameras, then requested + pixels, then contributors, the list structure the dense kernel returns. + """ + device = ids.jdata.device + inverse = sparse_tiles.inverse_indices + offsets_unique = ids.joffsets.to(device) + starts = offsets_unique[inverse] + counts = offsets_unique[1:][inverse] - starts + offsets = torch.zeros(counts.numel() + 1, dtype=torch.long, device=device) + offsets[1:] = counts.cumsum(0) + segment = torch.repeat_interleave(torch.arange(counts.numel(), device=device), counts) + within = torch.arange(int(offsets[-1].item()), device=device) - offsets[segment] + gather = starts[segment] + within + requested = sparse_tiles.pixels_to_render + # jidx is empty for a single camera, so derive each pixel's camera from the offsets instead. + requested_offsets = requested.joffsets.to(device).long() + pixels_per_camera = requested_offsets[1:] - requested_offsets[:-1] + num_cameras = pixels_per_camera.numel() + camera = torch.repeat_interleave(torch.arange(num_cameras, device=device), pixels_per_camera) + within_camera = torch.arange(camera.numel(), device=device) - requested_offsets[camera] + list_ids = torch.stack([camera, within_camera], dim=1).to(torch.int32) + + def nest(per_contribution: torch.Tensor) -> JaggedTensor: + data = per_contribution.index_select(0, gather) + expanded = JaggedTensor.from_data_offsets_and_list_ids(data, offsets, list_ids) + if len(expanded) == num_cameras: + return expanded + # from_data_offsets_and_list_ids takes the camera count from the largest list id, so cameras after + # the last one with a requested pixel are dropped (openvdb/fvdb-core#802). The nested constructor + # keeps them at the cost of one tensor per pixel on the host, so it is used only in that case. + per_pixel = JaggedTensor.from_data_and_offsets(data, offsets).unbind() + nested: list[list[torch.Tensor]] = [] + start = 0 + for count in pixels_per_camera.tolist(): + nested.append(list(per_pixel[start : start + count])) + start += count + return JaggedTensor(nested) + + return nest(ids.jdata), nest(weights.jdata) + + +def rasterize_contributing_gaussian_ids_sparse( + projected: ProjectedGaussians, + opacities: torch.Tensor, + sparse_tiles: SparseGaussianTileIntersection, + top_k_contributors: int = 0, +) -> tuple[JaggedTensor, JaggedTensor]: + """List contributing Gaussians, with blend weights, at the requested pixels only. + + Mode selection follows :func:`rasterize_contributing_gaussian_ids`. Results are returned in the + order of ``sparse_tiles.pixels_to_render``, duplicates included. + + Args: + projected (ProjectedGaussians): Output of :func:`project_gaussians`. + opacities (torch.Tensor): Per-camera opacities, ``[C, N]``, from :func:`compute_gaussian_opacities`. + sparse_tiles (SparseGaussianTileIntersection): Output of :func:`intersect_gaussian_tiles_sparse`. + top_k_contributors (int): Contributors to keep per pixel, or ``0`` for all of them. + + Returns: + gaussian_ids (JaggedTensor): Contributor indices, nested as cameras, then pixels, then contributors. + weights (JaggedTensor): Blend weight of each listed contributor, same structure. + """ + with torch.no_grad(): + opacities = check_opacities(opacities, projected) + check_tiles_match(sparse_tiles, projected) + num_contributing = None + if top_k_contributors <= 0: + num_contributing, _ = _count_sparse_unique(projected, opacities, sparse_tiles) + ids, weights = F.rasterize_contributing_gaussian_ids_sparse( + projected.means2d, + projected.conics, + opacities, + sparse_tiles.tile_offsets, + sparse_tiles.tile_gaussian_ids, + sparse_tiles.unique_pixels, + sparse_tiles.active_tiles, + sparse_tiles.tile_pixel_mask, + sparse_tiles.tile_pixel_cumsum, + sparse_tiles.pixel_map, + sparse_tiles.image_width, + sparse_tiles.image_height, + 0, + 0, + sparse_tiles.tile_size, + top_k_contributors, + num_contributing, + ) + if sparse_tiles.has_duplicates: + ids, weights = _expand_contributions(sparse_tiles, ids, weights) + return ids, weights diff --git a/fvdb_reality_capture/radiance_fields/_gaussian_autograd.py b/fvdb_reality_capture/functional/_autograd.py similarity index 56% rename from fvdb_reality_capture/radiance_fields/_gaussian_autograd.py rename to fvdb_reality_capture/functional/_autograd.py index 2c32479d..6b24486c 100644 --- a/fvdb_reality_capture/radiance_fields/_gaussian_autograd.py +++ b/fvdb_reality_capture/functional/_autograd.py @@ -1,18 +1,20 @@ # Copyright Contributors to the OpenVDB Project # SPDX-License-Identifier: Apache-2.0 # -"""Python torch.autograd.Function wrappers for Gaussian splatting dispatch functions. +"""``torch.autograd.Function`` wrappers over the Gaussian splatting kernels in :mod:`fvdb.functional`. + +``fvdb.functional`` exposes each kernel's forward and backward pass as separate, non-differentiable +functions. The classes here pair them up so that gradients flow through the composable stages in +this package. They are private; use the stage functions instead. """ from __future__ import annotations -from typing import Any, cast +from typing import Any import torch - -from fvdb import _fvdb_cpp as _C -from fvdb._fvdb_cpp import JaggedTensor as JaggedTensorCpp -from fvdb.jagged_tensor import JaggedTensor +from fvdb import JaggedTensor +from fvdb import functional as F # --------------------------------------------------------------------------- # Projection (analytic) @@ -20,7 +22,7 @@ class _ProjectGaussiansFn(torch.autograd.Function): - """Python autograd wrapper for the analytic Gaussian projection forward/backward dispatch.""" + """Analytic (EWA) projection of 3D Gaussians to 2D with gradients to the 3D parameters.""" @staticmethod def forward( @@ -42,7 +44,7 @@ def forward( accum_step_counts: torch.Tensor | None = None, accum_max_radii: torch.Tensor | None = None, ): - result = _C.project_gaussians_analytic_fwd( + radii, means2d, depths, conics, compensations = F.project_gaussians_analytic_fwd( means, quats, log_scales, @@ -57,11 +59,8 @@ def forward( calc_compensations, ortho, ) - radii: torch.Tensor = result[0] - means2d: torch.Tensor = result[1] - depths: torch.Tensor = result[2] - conics: torch.Tensor = result[3] - compensations: torch.Tensor | None = result[4] if calc_compensations else None + if not calc_compensations: + compensations = None to_save = [means, quats, log_scales, world_to_cam, projection_matrices, radii, conics] if compensations is not None: @@ -101,19 +100,13 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor grad_compensations = gc.contiguous() saved = ctx.saved_tensors - means = saved[0] - quats = saved[1] - log_scales = saved[2] - world_to_cam = saved[3] - projection_matrices = saved[4] - radii = saved[5] - conics = saved[6] + means, quats, log_scales, world_to_cam, projection_matrices, radii, conics = saved[:7] compensations = saved[7] if ctx.calc_compensations else None assert grad_means2d is not None assert grad_depths is not None assert grad_conics is not None - d_means, _, d_quats, d_scales, d_w2c = _C.project_gaussians_analytic_bwd( + d_means, _, d_quats, d_scales, d_w2c = F.project_gaussians_analytic_bwd( means, quats, log_scales, @@ -136,33 +129,16 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor ctx.accum_step_counts, ) - return ( - d_means, - d_quats, - d_scales, - d_w2c, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - ) + return (d_means, d_quats, d_scales, d_w2c) + (None,) * 12 # --------------------------------------------------------------------------- -# Projection (analytic, jagged) +# Projection (analytic, jagged batches of scenes) # --------------------------------------------------------------------------- class _ProjectGaussiansJaggedFn(torch.autograd.Function): - """Python autograd wrapper for the jagged Gaussian projection dispatch.""" + """Analytic projection for a batch of scenes with varying Gaussian and camera counts.""" @staticmethod def forward( @@ -182,7 +158,7 @@ def forward( min_radius_2d: float, ortho: bool, ): - result = _C.project_gaussians_analytic_jagged_fwd( + radii, means2d, depths, conics, compensations = F.project_gaussians_analytic_jagged_fwd( g_sizes, means, quats, @@ -198,25 +174,8 @@ def forward( min_radius_2d, ortho, ) - radii, means2d, depths, conics, compensations = ( - result[0], - result[1], - result[2], - result[3], - result[4], - ) - ctx.save_for_backward( - g_sizes, - means, - quats, - scales, - c_sizes, - world_to_cam, - projection_matrices, - radii, - conics, - ) + ctx.save_for_backward(g_sizes, means, quats, scales, c_sizes, world_to_cam, projection_matrices, radii, conics) ctx.image_width = image_width ctx.image_height = image_height ctx.eps2d = eps2d @@ -229,8 +188,7 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor grad_means2d = grad_outputs[1] grad_depths = grad_outputs[2] grad_conics = grad_outputs[3] - # grad_outputs[4] is grad_compensations -- the jagged backward dispatch - # does not consume it, so we ignore it here. + # grad_outputs[4] is grad_compensations, which the jagged backward does not consume. if grad_means2d is not None: grad_means2d = grad_means2d.contiguous() if grad_depths is not None: @@ -244,7 +202,7 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor assert grad_means2d is not None assert grad_depths is not None assert grad_conics is not None - d_means, _, d_quats, d_scales, d_w2c = _C.project_gaussians_analytic_jagged_bwd( + d_means, _, d_quats, d_scales, d_w2c = F.project_gaussians_analytic_jagged_bwd( g_sizes, means, quats, @@ -264,22 +222,7 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor ctx.ortho, ) - return ( - None, - d_means, - d_quats, - d_scales, - None, - d_w2c, - None, - None, - None, - None, - None, - None, - None, - None, - ) + return (None, d_means, d_quats, d_scales, None, d_w2c) + (None,) * 8 # --------------------------------------------------------------------------- @@ -288,7 +231,7 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor class _EvaluateGaussianSHFn(torch.autograd.Function): - """Python autograd wrapper for the SH evaluation forward/backward dispatch.""" + """Spherical-harmonics feature evaluation with gradients to the coefficients, means and cameras.""" @staticmethod def forward( @@ -297,13 +240,13 @@ def forward( num_cameras: int, means: torch.Tensor, # [N, 3] world_to_cam: torch.Tensor, # [C, 4, 4] - camera_ids: torch.Tensor, # empty (dense) or [nnz] int32 (jagged) - gaussian_ids: torch.Tensor, # empty (dense) or [nnz] int32 (jagged) - sh0_coeffs: torch.Tensor, # [N, 1, D] (dense) or [nnz, 1, D] (jagged) - shN_coeffs: torch.Tensor, # [N, K-1, D] (dense) or [nnz, K-1, D] (jagged) - radii: torch.Tensor, # [C, N, 2] (dense) or [1, nnz, 2] (jagged) + camera_ids: torch.Tensor, # empty (dense) or [M] int32 (packed) + gaussian_ids: torch.Tensor, # empty (dense) or [M] int32 (packed) + sh0_coeffs: torch.Tensor, # [N, 1, D] (dense) or [M, 1, D] (packed) + shN_coeffs: torch.Tensor, # [N, K-1, D] (dense) or [M, K-1, D] (packed) + radii: torch.Tensor, # [C, N, 2] (dense) or [1, M, 2] (packed) ) -> torch.Tensor: - render_quantities = _C.evaluate_spherical_harmonics_fwd( + features = F.evaluate_spherical_harmonics_fwd( sh_degree_to_use, num_cameras, means, @@ -315,30 +258,23 @@ def forward( radii, ) - ctx.save_for_backward( - means, - world_to_cam, - camera_ids, - gaussian_ids, - shN_coeffs, - radii, - ) + ctx.save_for_backward(means, world_to_cam, camera_ids, gaussian_ids, shN_coeffs, radii) ctx.sh_degree_to_use = sh_degree_to_use ctx.num_cameras = num_cameras ctx.num_gaussians = sh0_coeffs.size(0) - return render_quantities + return features @staticmethod def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor | None, ...]: d_loss_d_colors = grad_outputs[0] if d_loss_d_colors is None: - return (None, None, None, None, None, None, None, None, None) + return (None,) * 9 d_loss_d_colors = d_loss_d_colors.contiguous() means, world_to_cam, camera_ids, gaussian_ids, shN_coeffs, radii = ctx.saved_tensors - d_sh0, d_shN, d_means, d_w2c = _C.evaluate_spherical_harmonics_bwd( + d_sh0, d_shN, d_means, d_w2c = F.evaluate_spherical_harmonics_bwd( ctx.sh_degree_to_use, ctx.num_cameras, ctx.num_gaussians, @@ -361,15 +297,32 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor # --------------------------------------------------------------------------- +def _save_optional(ctx, to_save: list[torch.Tensor], backgrounds: torch.Tensor | None, masks: torch.Tensor | None): + """Append the optional rasterization inputs to ``to_save`` and record which were present.""" + ctx.has_backgrounds = backgrounds is not None + ctx.has_masks = masks is not None + if backgrounds is not None: + to_save.append(backgrounds) + if masks is not None: + to_save.append(masks) + + +def _load_optional(ctx, saved: tuple[torch.Tensor, ...], first: int) -> tuple[torch.Tensor | None, torch.Tensor | None]: + """Read back the optional rasterization inputs saved by :func:`_save_optional`.""" + backgrounds = saved[first] if ctx.has_backgrounds else None + masks = saved[first + int(ctx.has_backgrounds)] if ctx.has_masks else None + return backgrounds, masks + + class _RasterizeScreenSpaceGaussiansFn(torch.autograd.Function): - """Python autograd wrapper for the dense Gaussian rasterization forward/backward dispatch.""" + """Dense alpha-blending of projected Gaussians with gradients to the 2D quantities.""" @staticmethod def forward( ctx, means2d: torch.Tensor, conics: torch.Tensor, - colors: torch.Tensor, + features: torch.Tensor, opacities: torch.Tensor, image_width: int, image_height: int, @@ -382,10 +335,10 @@ def forward( backgrounds: torch.Tensor | None, masks: torch.Tensor | None, ): - result = _C.rasterize_screen_space_gaussians_fwd( + rendered, rendered_alphas, last_ids = F.rasterize_screen_space_gaussians_fwd( means2d, conics, - colors, + features, opacities, image_width, image_height, @@ -398,21 +351,9 @@ def forward( backgrounds, masks, ) - rendered_colors = result[0] - rendered_alphas = result[1] - last_ids = result[2] - - to_save = [means2d, conics, colors, opacities, tile_offsets, tile_gaussian_ids, rendered_alphas, last_ids] - if backgrounds is not None: - to_save.append(backgrounds) - ctx.has_backgrounds = True - else: - ctx.has_backgrounds = False - if masks is not None: - to_save.append(masks) - ctx.has_masks = True - else: - ctx.has_masks = False + + to_save = [means2d, conics, features, opacities, tile_offsets, tile_gaussian_ids, rendered_alphas, last_ids] + _save_optional(ctx, to_save, backgrounds, masks) ctx.save_for_backward(*to_save) ctx.image_width = image_width @@ -422,37 +363,26 @@ def forward( ctx.tile_size = tile_size ctx.absgrad = absgrad - return rendered_colors, rendered_alphas + return rendered, rendered_alphas @staticmethod def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor | None, ...]: - d_loss_d_rendered_colors = grad_outputs[0] - d_loss_d_rendered_alphas = grad_outputs[1] - if d_loss_d_rendered_colors is not None: - d_loss_d_rendered_colors = d_loss_d_rendered_colors.contiguous() - if d_loss_d_rendered_alphas is not None: - d_loss_d_rendered_alphas = d_loss_d_rendered_alphas.contiguous() + d_rendered, d_alphas = grad_outputs[0], grad_outputs[1] + if d_rendered is not None: + d_rendered = d_rendered.contiguous() + if d_alphas is not None: + d_alphas = d_alphas.contiguous() saved = ctx.saved_tensors - means2d, conics, colors, opacities = saved[0], saved[1], saved[2], saved[3] - tile_offsets, tile_gaussian_ids = saved[4], saved[5] - rendered_alphas, last_ids = saved[6], saved[7] - - backgrounds: torch.Tensor | None = None - masks: torch.Tensor | None = None - opt_idx = 8 - if ctx.has_backgrounds: - backgrounds = saved[opt_idx] - opt_idx += 1 - if ctx.has_masks: - masks = saved[opt_idx] - - assert d_loss_d_rendered_colors is not None - assert d_loss_d_rendered_alphas is not None - result = _C.rasterize_screen_space_gaussians_bwd( + means2d, conics, features, opacities, tile_offsets, tile_gaussian_ids, rendered_alphas, last_ids = saved[:8] + backgrounds, masks = _load_optional(ctx, saved, 8) + + 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_bwd( means2d, conics, - colors, + features, opacities, ctx.image_width, ctx.image_height, @@ -463,34 +393,15 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor tile_gaussian_ids, rendered_alphas, last_ids, - d_loss_d_rendered_colors, - d_loss_d_rendered_alphas, + d_rendered, + d_alphas, ctx.absgrad, -1, backgrounds, masks, ) - d_means2d = result[1] - d_conics = result[2] - d_colors = result[3] - d_opacities = result[4] - - return ( - d_means2d, - d_conics, - d_colors, - d_opacities, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - ) + + return (d_means2d, d_conics, d_features, d_opacities) + (None,) * 10 # --------------------------------------------------------------------------- @@ -499,7 +410,11 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor class _RasterizeScreenSpaceGaussiansSparseFn(torch.autograd.Function): - """Python autograd wrapper for the sparse Gaussian rasterization forward/backward dispatch.""" + """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. + """ @staticmethod def forward( @@ -524,8 +439,8 @@ def forward( backgrounds: torch.Tensor | None, masks: torch.Tensor | None, ): - result = _C.rasterize_screen_space_gaussians_sparse_fwd( - pixels_to_render._impl, + rendered_jt, alphas_jt, last_ids_jt = F.rasterize_screen_space_gaussians_sparse_fwd( + pixels_to_render, means2d, conics, features, @@ -545,13 +460,6 @@ def forward( backgrounds, masks, ) - rendered_colors_jt = JaggedTensor(impl=result[0]) - rendered_alphas_jt = JaggedTensor(impl=result[1]) - last_ids_jt = JaggedTensor(impl=result[2]) - - joffsets = pixels_to_render.joffsets - jidx = pixels_to_render.jidx - jlidx = pixels_to_render.jlidx to_save = [ means2d, @@ -561,27 +469,16 @@ def forward( tile_offsets, tile_gaussian_ids, pixels_to_render.jdata, - rendered_colors_jt.jdata, - rendered_alphas_jt.jdata, + pixels_to_render.joffsets, + pixels_to_render.jlidx, + alphas_jt.jdata, last_ids_jt.jdata, - joffsets, - jidx, - jlidx, active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map, ] - if backgrounds is not None: - to_save.append(backgrounds) - ctx.has_backgrounds = True - else: - ctx.has_backgrounds = False - if masks is not None: - to_save.append(masks) - ctx.has_masks = True - else: - ctx.has_masks = False + _save_optional(ctx, to_save, backgrounds, masks) ctx.save_for_backward(*to_save) ctx.image_width = image_width @@ -590,48 +487,28 @@ def forward( ctx.image_origin_h = image_origin_h ctx.tile_size = tile_size ctx.absgrad = absgrad - ctx.num_outer_lists = len(pixels_to_render) - return rendered_colors_jt.jdata, rendered_alphas_jt.jdata + return rendered_jt.jdata, alphas_jt.jdata @staticmethod def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor | None, ...]: - d_loss_d_rendered_features_jdata = grad_outputs[0] - d_loss_d_rendered_alphas_jdata = grad_outputs[1] - if d_loss_d_rendered_features_jdata is not None: - d_loss_d_rendered_features_jdata = d_loss_d_rendered_features_jdata.contiguous() - if d_loss_d_rendered_alphas_jdata is not None: - d_loss_d_rendered_alphas_jdata = d_loss_d_rendered_alphas_jdata.contiguous() + d_rendered, d_alphas = grad_outputs[0], grad_outputs[1] + if d_rendered is not None: + d_rendered = d_rendered.contiguous() + if d_alphas is not None: + d_alphas = d_alphas.contiguous() saved = ctx.saved_tensors - means2d, conics, features, opacities = saved[0], saved[1], saved[2], saved[3] - tile_offsets, tile_gaussian_ids = saved[4], saved[5] - pixels_jdata = saved[6] - rendered_alphas_jdata = saved[8] - last_ids_jdata = saved[9] - joffsets, jidx, jlidx = saved[10], saved[11], saved[12] - active_tiles = saved[13] - tile_pixel_mask, tile_pixel_cumsum, pixel_map = saved[14], saved[15], saved[16] - - backgrounds: torch.Tensor | None = None - masks: torch.Tensor | None = None - opt_idx = 17 - if ctx.has_backgrounds: - backgrounds = saved[opt_idx] - opt_idx += 1 - if ctx.has_masks: - masks = saved[opt_idx] - - pixels_jt = JaggedTensor(impl=_C.JaggedTensor.from_data_offsets_and_list_ids(pixels_jdata, joffsets, jlidx)) - rendered_alphas_jt = pixels_jt.jagged_like(rendered_alphas_jdata) - last_ids_jt = pixels_jt.jagged_like(last_ids_jdata) - assert d_loss_d_rendered_features_jdata is not None - assert d_loss_d_rendered_alphas_jdata is not None - d_loss_d_rendered_features_jt = pixels_jt.jagged_like(d_loss_d_rendered_features_jdata) - d_loss_d_rendered_alphas_jt = pixels_jt.jagged_like(d_loss_d_rendered_alphas_jdata) - - result = _C.rasterize_screen_space_gaussians_sparse_bwd( - pixels_jt._impl, + 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) + + pixels_jt = JaggedTensor.from_data_offsets_and_list_ids(pixels_jdata, joffsets, jlidx) + 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( + pixels_jt, means2d, conics, features, @@ -643,10 +520,10 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor ctx.tile_size, tile_offsets, tile_gaussian_ids, - rendered_alphas_jt._impl, - last_ids_jt._impl, - d_loss_d_rendered_features_jt._impl, - d_loss_d_rendered_alphas_jt._impl, + pixels_jt.jagged_like(alphas_jdata), + pixels_jt.jagged_like(last_ids_jdata), + pixels_jt.jagged_like(d_rendered), + pixels_jt.jagged_like(d_alphas), active_tiles, tile_pixel_mask, tile_pixel_cumsum, @@ -656,32 +533,8 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor backgrounds, masks, ) - d_means2d = result[1] - d_conics = result[2] - d_colors = result[3] - d_opacities = result[4] - - return ( - d_means2d, - d_conics, - d_colors, - d_opacities, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - ) + + return (d_means2d, d_conics, d_features, d_opacities) + (None,) * 15 # --------------------------------------------------------------------------- @@ -690,7 +543,7 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor class _RasterizeWorldSpaceGaussiansFn(torch.autograd.Function): - """Python autograd wrapper for world-space Gaussian rasterization forward/backward dispatch.""" + """Ray-based rasterization of 3D Gaussians with gradients to the 3D parameters.""" @staticmethod def forward( @@ -716,7 +569,7 @@ def forward( backgrounds: torch.Tensor | None, masks: torch.Tensor | None, ): - result = _C.rasterize_world_space_gaussians_fwd( + rendered, rendered_alphas, last_ids = F.rasterize_world_space_gaussians_fwd( means, quats, log_scales, @@ -726,8 +579,8 @@ def forward( world_to_cam_end, projection_matrices, distortion_coeffs, - _C.RollingShutterType(rolling_shutter_type), - _C.CameraModel(camera_model), + rolling_shutter_type, + camera_model, image_width, image_height, image_origin_w, @@ -738,9 +591,6 @@ def forward( backgrounds, masks, ) - rendered_features = result[0] - rendered_alphas = result[1] - last_ids = result[2] to_save = [ means, @@ -757,16 +607,7 @@ def forward( rendered_alphas, last_ids, ] - if backgrounds is not None: - to_save.append(backgrounds) - ctx.has_backgrounds = True - else: - ctx.has_backgrounds = False - if masks is not None: - to_save.append(masks) - ctx.has_masks = True - else: - ctx.has_masks = False + _save_optional(ctx, to_save, backgrounds, masks) ctx.save_for_backward(*to_save) ctx.image_width = image_width @@ -777,37 +618,25 @@ def forward( ctx.rolling_shutter_type = rolling_shutter_type ctx.camera_model = camera_model - return rendered_features, rendered_alphas + return rendered, rendered_alphas @staticmethod def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor | None, ...]: - d_loss_d_rendered_features = grad_outputs[0] - d_loss_d_rendered_alphas = grad_outputs[1] - if d_loss_d_rendered_features is not None: - d_loss_d_rendered_features = d_loss_d_rendered_features.contiguous() - if d_loss_d_rendered_alphas is not None: - d_loss_d_rendered_alphas = d_loss_d_rendered_alphas.contiguous() + d_rendered, d_alphas = grad_outputs[0], grad_outputs[1] + if d_rendered is not None: + d_rendered = d_rendered.contiguous() + if d_alphas is not None: + d_alphas = d_alphas.contiguous() saved = ctx.saved_tensors - means, quats, log_scales = saved[0], saved[1], saved[2] - features, opacities = saved[3], saved[4] - world_to_cam_start, world_to_cam_end = saved[5], saved[6] - projection_matrices, distortion_coeffs = saved[7], saved[8] - tile_offsets, tile_gaussian_ids = saved[9], saved[10] - rendered_alphas, last_ids = saved[11], saved[12] - - backgrounds: torch.Tensor | None = None - masks: torch.Tensor | None = None - opt_idx = 13 - if ctx.has_backgrounds: - backgrounds = saved[opt_idx] - opt_idx += 1 - if ctx.has_masks: - masks = saved[opt_idx] - - assert d_loss_d_rendered_features is not None - assert d_loss_d_rendered_alphas is not None - result = _C.rasterize_world_space_gaussians_bwd( + means, quats, log_scales, features, opacities = saved[:5] + world_to_cam_start, world_to_cam_end, projection_matrices, distortion_coeffs = saved[5:9] + tile_offsets, tile_gaussian_ids, rendered_alphas, last_ids = saved[9:13] + backgrounds, masks = _load_optional(ctx, saved, 13) + + assert d_rendered is not None + assert d_alphas is not None + d_means, d_quats, d_log_scales, d_features, d_opacities = F.rasterize_world_space_gaussians_bwd( means, quats, log_scales, @@ -817,8 +646,8 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor world_to_cam_end, projection_matrices, distortion_coeffs, - _C.RollingShutterType(ctx.rolling_shutter_type), - _C.CameraModel(ctx.camera_model), + ctx.rolling_shutter_type, + ctx.camera_model, ctx.image_width, ctx.image_height, ctx.image_origin_w, @@ -828,36 +657,10 @@ def backward(ctx: Any, *grad_outputs: torch.Tensor | None) -> tuple[torch.Tensor tile_gaussian_ids, rendered_alphas, last_ids, - d_loss_d_rendered_features, - d_loss_d_rendered_alphas, + d_rendered, + d_alphas, backgrounds, masks, ) - d_means = result[0] - d_quats = result[1] - d_log_scales = result[2] - d_features = result[3] - d_opacities = result[4] - - return ( - d_means, - d_quats, - d_log_scales, - d_features, - d_opacities, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - None, - ) + + return (d_means, d_quats, d_log_scales, d_features, d_opacities) + (None,) * 15 diff --git a/fvdb_reality_capture/functional/_opacity.py b/fvdb_reality_capture/functional/_opacity.py new file mode 100644 index 00000000..705757f0 --- /dev/null +++ b/fvdb_reality_capture/functional/_opacity.py @@ -0,0 +1,51 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +"""Per-camera opacity computation shared by the rasterization and analysis stages.""" + +from __future__ import annotations + +import torch + +from ._types import ProjectedGaussians + + +def compute_gaussian_opacities(logit_opacities: torch.Tensor, projected: ProjectedGaussians) -> torch.Tensor: + """Turn logit opacities into the per-camera opacities the rasterization kernels consume. + + Applies the sigmoid, repeats the result for every camera, and multiplies in the anti-aliasing + compensation factors when the projection computed them. + + Args: + logit_opacities (torch.Tensor): Logit opacities, shape ``[N]``. + projected (ProjectedGaussians): Projection whose camera count and compensations to use. + + Returns: + opacities (torch.Tensor): Opacities in ``[0, 1]``, shape ``[C, N]``, contiguous. + """ + # The kernels require a contiguous [C, N] tensor, so the per-camera copy is materialized. + opacities = torch.sigmoid(logit_opacities).repeat(projected.num_cameras, 1) + if projected.compensations is not None: + opacities = opacities * projected.compensations + return opacities + + +def check_opacities(opacities: torch.Tensor, projected: ProjectedGaussians) -> torch.Tensor: + """Validate per-camera opacities against a projection and return them ready for the kernels. + + Args: + opacities (torch.Tensor): Opacities from :func:`compute_gaussian_opacities`, shape ``[C, N]``. + projected (ProjectedGaussians): Projection the opacities were computed for. + + Returns: + opacities (torch.Tensor): The same values, contiguous and on the projection's device. + """ + if tuple(opacities.shape) != (projected.num_cameras, projected.num_gaussians): + raise ValueError( + f"opacities must have shape [{projected.num_cameras}, {projected.num_gaussians}], got {tuple(opacities.shape)}" + ) + if opacities.device != projected.means2d.device: + raise ValueError(f"opacities must be on {projected.means2d.device}, got {opacities.device}") + # The kernels index a dense [C, N] buffer, so an expanded or strided view is materialized here + # rather than failing inside the kernel. + return opacities.contiguous() diff --git a/fvdb_reality_capture/functional/_projection.py b/fvdb_reality_capture/functional/_projection.py new file mode 100644 index 00000000..f17c73c1 --- /dev/null +++ b/fvdb_reality_capture/functional/_projection.py @@ -0,0 +1,214 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +"""Stage 1 of the composable pipeline: project 3D Gaussians into every camera.""" + +from __future__ import annotations + +from typing import cast + +import torch +from fvdb import functional as F + +from ..enums import CameraModel, ProjectionMethod +from ._autograd import _ProjectGaussiansFn +from ._types import ProjectedGaussians + + +def requires_distortion_coeffs(camera_model: CameraModel) -> bool: + """Whether a camera model is a lens-distortion model. + + Every model other than pinhole and orthographic distorts, needs distortion coefficients, and can + only be projected with the unscented transform. Keeping this the single test means a distortion + model added to fvdb is handled consistently by projection, rasterization and :class:`~fvdb_reality_capture.GaussianSplat3d`. + + Args: + camera_model (CameraModel): The camera model. + + Returns: + distorted (bool): ``True`` for the distortion models, ``False`` for pinhole and orthographic. + """ + return CameraModel(camera_model) not in (CameraModel.PINHOLE, CameraModel.ORTHOGRAPHIC) + + +def check_distortion_coeffs( + distortion_coeffs: torch.Tensor | None, camera_model: CameraModel, num_cameras: int, device: torch.device +) -> torch.Tensor | None: + """Check packed distortion coefficients for a camera batch before a kernel reads them. + + Pinhole and orthographic cameras ignore the coefficients, so for them ``None`` is returned whatever was + passed. For the distortion models the kernels read twelve coefficients per camera, so anything else + is rejected here rather than rendering wrong pixels or tripping a device assert. + + Args: + distortion_coeffs (torch.Tensor | None): Packed coefficients, ``[C, 12]``, or ``None``. + camera_model (CameraModel): The batch's camera model; distortion models require the coefficients. + num_cameras (int): ``C``. + device (torch.device): Device of the Gaussians and cameras. + + Returns: + distortion_coeffs (torch.Tensor | None): The coefficients the kernel should receive: ``None`` for + camera models without distortion, the validated tensor otherwise. + + Raises: + RuntimeError: If a distortion camera model has no coefficients, or its tensor is not a contiguous + ``[C, 12]`` tensor on ``device``. + """ + if not requires_distortion_coeffs(camera_model): + return None + if distortion_coeffs is None: + raise RuntimeError("distortionCoeffs must be provided for OpenCV camera models") + if list(distortion_coeffs.shape) != [num_cameras, 12]: + raise RuntimeError(f"distortionCoeffs must have shape ({num_cameras}, 12)") + if not distortion_coeffs.is_contiguous(): + raise RuntimeError("distortionCoeffs must be contiguous") + if distortion_coeffs.device != device: + raise RuntimeError(f"distortionCoeffs must be on {device}, got {distortion_coeffs.device}") + return distortion_coeffs + + +def resolve_projection_method(camera_model: CameraModel, projection_method: ProjectionMethod) -> ProjectionMethod: + """Replace :attr:`~fvdb_reality_capture.ProjectionMethod.AUTO` with the concrete method for a camera model. + + Pinhole and orthographic cameras default to the analytic projection; the distortion models default + to the unscented transform, which is the only method that supports them. + + Args: + camera_model (CameraModel): The camera model. + projection_method (ProjectionMethod): The requested method, possibly ``AUTO``. + + Returns: + projection_method (ProjectionMethod): ``ANALYTIC`` or ``UNSCENTED``. + """ + if projection_method != ProjectionMethod.AUTO: + return ProjectionMethod(projection_method) + if requires_distortion_coeffs(camera_model): + return ProjectionMethod.UNSCENTED + return ProjectionMethod.ANALYTIC + + +def project_gaussians( + means: torch.Tensor, + quats: torch.Tensor, + log_scales: torch.Tensor, + world_to_camera_matrices: torch.Tensor, + projection_matrices: torch.Tensor, + image_width: int, + image_height: int, + near: float = 0.01, + far: float = 1e10, + camera_model: CameraModel = CameraModel.PINHOLE, + projection_method: ProjectionMethod = ProjectionMethod.AUTO, + distortion_coeffs: torch.Tensor | None = None, + min_radius_2d: float = 0.0, + eps_2d: float = 0.3, + antialias: bool = False, + *, + accumulated_mean_2d_gradient_norms: torch.Tensor | None = None, + accumulated_gradient_step_counts: torch.Tensor | None = None, + accumulated_max_2d_radii: torch.Tensor | None = None, +) -> ProjectedGaussians: + """Project 3D Gaussians onto the image planes of a set of cameras. + + With the analytic method the result is differentiable with respect to ``means``, ``quats``, + ``log_scales`` and ``world_to_camera_matrices``. The unscented transform has no backward pass; + to train through it, rasterize with :func:`rasterize_world_space_gaussians`, which + differentiates through the 3D parameters directly. + + Args: + means (torch.Tensor): Gaussian centers in world space, ``[N, 3]``. + quats (torch.Tensor): Gaussian rotations as quaternions, ``[N, 4]``. + log_scales (torch.Tensor): Natural-log scale factors, ``[N, 3]``. + world_to_camera_matrices (torch.Tensor): World-to-camera transforms, ``[C, 4, 4]``, contiguous. + projection_matrices (torch.Tensor): Camera intrinsics, ``[C, 3, 3]``, contiguous. + image_width (int): Image width in pixels. + image_height (int): Image height in pixels. + near (float): Near clipping plane. Gaussians closer than this are culled. + far (float): Far clipping plane. Gaussians farther than this are culled. + camera_model (CameraModel): Camera model of all ``C`` cameras. + projection_method (ProjectionMethod): Projection method; ``AUTO`` picks per camera model. + distortion_coeffs (torch.Tensor | None): Packed OpenCV distortion coefficients, ``[C, 12]``, + contiguous. Required for the OpenCV camera models and ignored otherwise. + min_radius_2d (float): Gaussians whose projected radius is at most this many pixels are culled. + eps_2d (float): Blur added to the projected covariance for numerical stability. + antialias (bool): Compute opacity compensation factors for the blur added by ``eps_2d``. + accumulated_mean_2d_gradient_norms (torch.Tensor | None): Optional ``[N]`` float accumulator that + the analytic backward pass adds image-normalized 2D mean gradient norms into. + accumulated_gradient_step_counts (torch.Tensor | None): Optional ``[N]`` ``int32`` accumulator of + backward passes per Gaussian. Updated only when given together with the gradient-norm accumulator. + accumulated_max_2d_radii (torch.Tensor | None): Optional ``[N]`` ``int32`` accumulator of the + largest projected radius seen per Gaussian. Only updated alongside the other two. + + Returns: + projected (ProjectedGaussians): The projected Gaussians. + """ + if not projection_matrices.is_contiguous(): + raise RuntimeError("projectionMatrices must be contiguous") + if not world_to_camera_matrices.is_contiguous(): + raise RuntimeError("worldToCameraMatrices must be contiguous") + camera_model = CameraModel(camera_model) + num_cameras = world_to_camera_matrices.size(0) + distortion_coeffs = check_distortion_coeffs(distortion_coeffs, camera_model, num_cameras, means.device) + + resolved = resolve_projection_method(camera_model, projection_method) + if requires_distortion_coeffs(camera_model) and resolved != ProjectionMethod.UNSCENTED: + raise RuntimeError("OpenCV camera models require ProjectionMethod::UNSCENTED or AUTO") + + if resolved == ProjectionMethod.UNSCENTED: + if distortion_coeffs is None: + distortion_coeffs = torch.empty(num_cameras, 0, device=means.device, dtype=means.dtype) + radii, means2d, depths, conics, compensations = F.project_gaussians_ut_fwd( + means, + quats, + log_scales, + world_to_camera_matrices, + world_to_camera_matrices, + projection_matrices, + distortion_coeffs, + camera_model, + image_width, + image_height, + eps_2d, + near, + far, + min_radius_2d, + antialias, + ) + if not antialias: + compensations = None + else: + result = cast( + tuple[torch.Tensor, ...], + _ProjectGaussiansFn.apply( + means, + quats, + log_scales, + world_to_camera_matrices, + projection_matrices, + image_width, + image_height, + eps_2d, + near, + far, + min_radius_2d, + antialias, + camera_model == CameraModel.ORTHOGRAPHIC, + accumulated_mean_2d_gradient_norms, + accumulated_gradient_step_counts, + accumulated_max_2d_radii, + ), + ) + radii, means2d, depths, conics = result[:4] + compensations = result[4] if antialias else None + + return ProjectedGaussians( + radii=radii, + means2d=means2d, + depths=depths, + conics=conics, + compensations=compensations, + image_width=image_width, + image_height=image_height, + camera_model=camera_model, + projection_method=resolved, + ) diff --git a/fvdb_reality_capture/functional/_rasterization.py b/fvdb_reality_capture/functional/_rasterization.py new file mode 100644 index 00000000..85958ed4 --- /dev/null +++ b/fvdb_reality_capture/functional/_rasterization.py @@ -0,0 +1,467 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +"""Stage 4 of the composable pipeline: alpha-blend features into images or pixel sets.""" + +from __future__ import annotations + +from typing import cast + +import torch +import torch.nn.functional as nnf +from fvdb import JaggedTensor + +from ..enums import RollingShutterType +from ._autograd import ( + _RasterizeScreenSpaceGaussiansFn, + _RasterizeScreenSpaceGaussiansSparseFn, + _RasterizeWorldSpaceGaussiansFn, +) +from ._opacity import check_opacities +from ._projection import check_distortion_coeffs +from ._tile_intersection import check_tiles_match +from ._types import GaussianTileIntersection, ProjectedGaussians, SparseGaussianTileIntersection + +Crop = tuple[int, int, int, int] +"""A crop as ``(origin_w, origin_h, width, height)`` in pixels.""" + + +def validate_crop(crop: Crop, image_width: int, image_height: int) -> Crop: + """Check a crop against an image and clip it to the image bounds. + + A crop that runs past the image edge is clipped; one that lies entirely outside the image clips to a + zero-size crop, so rendering it yields an empty ``[C, 0, 0, D]`` result rather than an error, and + :func:`pad_crop` can grow that back to the requested size. + + Args: + crop (tuple[int, int, int, int]): ``(origin_w, origin_h, width, height)`` in pixels. + image_width (int): Image width in pixels. + image_height (int): Image height in pixels. + + Returns: + crop (tuple[int, int, int, int]): The crop with its size clipped so it lies inside the image. + + Raises: + ValueError: If the origin is negative or the size is not positive. + """ + origin_w, origin_h, width, height = crop + if origin_w < 0 or origin_h < 0: + raise ValueError(f"Crop origin must be non-negative, got ({origin_w}, {origin_h})") + if width <= 0 or height <= 0: + raise ValueError(f"Crop size must be positive, got ({width}, {height})") + width = max(0, min(width, image_width - origin_w)) + height = max(0, min(height, image_height - origin_h)) + if width == 0 or height == 0: + return origin_w, origin_h, 0, 0 + return origin_w, origin_h, width, height + + +def _empty_render(features: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """The ``[C, 0, 0, D]`` render of a crop clipped to nothing, as a slice of ``features`` so it stays connected.""" + empty = features[:, :0].reshape(features.shape[0], 0, 0, features.shape[-1]) + return empty, empty[..., :1] + + +def _finish_dense_render( + images: torch.Tensor, + alphas: torch.Tensor, + crop: Crop | None, + masks: torch.Tensor | None, + backgrounds: torch.Tensor | None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Slice a full-size render to its crop, then apply the per-pixel mask, already in crop coordinates.""" + images, alphas = apply_crop(images, alphas, crop) + if masks is not None: + images, alphas = apply_pixel_mask(images, alphas, masks, backgrounds) + return images, alphas + + +def _render_masks( + crop: Crop | None, + masks: torch.Tensor | None, + crop_masks: torch.Tensor | None, + projected: ProjectedGaussians, + tiles: GaussianTileIntersection, + device: torch.device, +) -> tuple[Crop | None, torch.Tensor | None, torch.Tensor | None]: + """Resolve the crop and masks into what the rasterizer and the post-pass need. + + Checks that ``tiles`` belong to ``projected`` first. Returns the clipped crop, the per-pixel mask in + the output's coordinates to apply after rendering and slicing (``None`` when no pixel mask was given, + since slicing to the crop already discards out-of-crop pixels), and the per-tile mask that lets the + rasterizer skip tiles outside the crop or fully masked out. ``masks`` is always in image coordinates + and ``crop_masks`` always in crop coordinates (the crop as requested or as clipped to the image), so + a shape never has to be guessed; the tile mask is pooled from whichever was given, without building a + full-image mask. The tile grid comes from ``tiles`` so it cannot drift from the intersection's. + """ + check_tiles_match(tiles, projected) + num_cameras = projected.num_cameras + tile_size = tiles.tile_size + full_shape = (num_cameras, tiles.image_height, tiles.image_width) + if masks is not None and crop_masks is not None: + raise ValueError("pass either masks (image coordinates) or crop_masks (crop coordinates), not both") + if crop_masks is not None and crop is None: + raise ValueError("crop_masks needs a crop") + for name, tensor in (("masks", masks), ("crop_masks", crop_masks)): + if tensor is not None and tensor.device != device: + raise ValueError(f"{name} must be on {device}, got {tensor.device}") + if masks is not None and tuple(masks.shape) != full_shape: + raise ValueError(f"masks must be a full-image [C, H, W] mask of shape {full_shape}, got {tuple(masks.shape)}") + if crop is None: + if masks is None: + return None, None, None + masks = masks.bool() + return None, masks, pixel_mask_to_tile_mask(masks, tile_size) + requested_shape = (num_cameras, crop[3], crop[2]) + origin_w, origin_h, width, height = clipped = validate_crop(crop, tiles.image_width, tiles.image_height) + tile_y0, tile_x0 = origin_h // tile_size, origin_w // tile_size + tile_y1, tile_x1 = -(-(origin_h + height) // tile_size), -(-(origin_w + width) // tile_size) + tile_window = torch.zeros(num_cameras, tiles.num_tiles_h, tiles.num_tiles_w, dtype=torch.bool, device=device) + tile_window[:, tile_y0:tile_y1, tile_x0:tile_x1] = True + if masks is None and crop_masks is None: + return clipped, None, tile_window + crop_shape = (num_cameras, height, width) + if masks is not None: + crop_mask = _window(masks.bool(), clipped) + elif tuple(crop_masks.shape) in (crop_shape, requested_shape): + # A mask of the requested crop size covers the clipped part in its top-left corner. + crop_mask = crop_masks[:, :height, :width].bool() + else: + raise ValueError( + f"crop_masks must match the crop {requested_shape} or its clipped size {crop_shape}, got {tuple(crop_masks.shape)}" + ) + # Pool the crop mask over the crop's tiles. The crop need not start on a tile boundary, so the mask is + # placed at its offset within the first tile before pooling. + aligned = torch.zeros( + num_cameras, origin_h % tile_size + height, origin_w % tile_size + width, dtype=torch.bool, device=device + ) + aligned[:, origin_h % tile_size :, origin_w % tile_size :] = crop_mask + tile_masks = torch.zeros_like(tile_window) + if height > 0 and width > 0: + tile_masks[:, tile_y0:tile_y1, tile_x0:tile_x1] = pixel_mask_to_tile_mask(aligned, tile_size) + return clipped, crop_mask, tile_masks & tile_window + + +def apply_crop(images: torch.Tensor, alphas: torch.Tensor, crop: Crop | None) -> tuple[torch.Tensor, torch.Tensor]: + """Slice rendered images and alphas to a crop window. + + Args: + images (torch.Tensor): Rendered images, ``[C, H, W, D]``. + alphas (torch.Tensor): Rendered alphas, ``[C, H, W, 1]``. + crop (tuple[int, int, int, int] | None): ``(origin_w, origin_h, width, height)`` window, already + validated against the image; ``None`` returns the inputs unchanged. + + Returns: + images (torch.Tensor): The window of ``images``, ``[C, height, width, D]``. + alphas (torch.Tensor): The window of ``alphas``, ``[C, height, width, 1]``. + """ + if crop is None: + return images, alphas + return _window(images, crop), _window(alphas, crop) + + +def _window(tensor: torch.Tensor, crop: Crop) -> torch.Tensor: + origin_w, origin_h, width, height = crop + return tensor[:, origin_h : origin_h + height, origin_w : origin_w + width] + + +def _background_like(images: torch.Tensor, backgrounds: torch.Tensor | None) -> torch.Tensor: + """Per-camera background features as ``[C, 1, 1, D]``, broadcastable against ``images``; black if ``None``.""" + if backgrounds is None: + return torch.zeros(images.shape[0], 1, 1, images.shape[-1], device=images.device, dtype=images.dtype) + return backgrounds.to(images)[:, None, None, :] + + +def pad_crop( + images: torch.Tensor, + alphas: torch.Tensor, + height: int, + width: int, + backgrounds: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Extend a rendered crop to ``height`` by ``width``, filling the added area with the background at zero alpha. + + This is how a crop that runs past the image edge keeps the size that was asked for: the part inside + the image is rendered, the rest is background. Inputs already of the target size are returned as is. + + Args: + images (torch.Tensor): Rendered crop, ``[C, h, w, D]`` with ``h <= height`` and ``w <= width``. + alphas (torch.Tensor): Its alphas, ``[C, h, w, 1]``. + height (int): Target height in pixels. + width (int): Target width in pixels. + backgrounds (torch.Tensor | None): Per-camera background features, ``[C, D]``. Black if ``None``. + + Returns: + images (torch.Tensor): ``[C, height, width, D]`` with the input in its top-left corner. + alphas (torch.Tensor): ``[C, height, width, 1]``, zero outside the input. + + Raises: + ValueError: If the input is larger than the target size in either dimension. + """ + num_cameras, current_h, current_w, channels = images.shape + if (current_h, current_w) == (height, width): + return images, alphas + if current_h > height or current_w > width: + raise ValueError(f"pad_crop cannot shrink a {(current_h, current_w)} render to {(height, width)}") + padded = _background_like(images, backgrounds).expand(num_cameras, height, width, channels).clone() + padded[:, :current_h, :current_w] = images + padded_alphas = alphas.new_zeros(num_cameras, height, width, 1) + padded_alphas[:, :current_h, :current_w] = alphas + return padded, padded_alphas + + +def pixel_mask_to_tile_mask(pixel_mask: torch.Tensor, tile_size: int) -> torch.Tensor: + """Mark a tile as rendered if any of its pixels is. + + Args: + pixel_mask (torch.Tensor): Boolean per-pixel mask, ``[C, H, W]``. + tile_size (int): Tile side length in pixels. + + Returns: + tile_mask (torch.Tensor): Boolean per-tile mask, ``[C, ceil(H / tile_size), ceil(W / tile_size)]``. + """ + pooled = nnf.max_pool2d( + pixel_mask.bool().unsqueeze(1).float(), kernel_size=tile_size, stride=tile_size, ceil_mode=True + ) + return pooled.bool().squeeze(1) + + +def apply_pixel_mask( + images: torch.Tensor, alphas: torch.Tensor, pixel_mask: torch.Tensor, backgrounds: torch.Tensor | None +) -> tuple[torch.Tensor, torch.Tensor]: + """Fill masked-out pixels of a render with the background at zero alpha. + + The rasterizers skip whole tiles; this is the per-pixel pass that follows, and it works on any + render, so a crop can be masked after slicing without building a full-image mask. + + Args: + images (torch.Tensor): Rendered features, ``[C, H, W, D]``. + alphas (torch.Tensor): Rendered alphas, ``[C, H, W, 1]``. + pixel_mask (torch.Tensor): Boolean mask, ``[C, H, W]``; ``True`` keeps the rendered pixel. + backgrounds (torch.Tensor | None): Per-camera background features, ``[C, D]``. Black if ``None``. + + Returns: + images (torch.Tensor): ``[C, H, W, D]`` with masked-out pixels set to the background. + alphas (torch.Tensor): ``[C, H, W, 1]`` with masked-out pixels set to zero. + """ + keep = pixel_mask.unsqueeze(-1).to(images.dtype) + background = _background_like(images, backgrounds) + return images * keep + background * (1.0 - keep), alphas * keep + + +def rasterize_screen_space_gaussians( + projected: ProjectedGaussians, + features: torch.Tensor, + opacities: torch.Tensor, + tiles: GaussianTileIntersection, + backgrounds: torch.Tensor | None = None, + masks: torch.Tensor | None = None, + crop: Crop | None = None, + crop_masks: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Alpha-blend projected Gaussians into dense images. + + Differentiable with respect to ``features``, ``opacities`` and, for an ``ANALYTIC`` projection, + the projection itself. The unscented projection is forward-only, so through this function the 3D + parameters receive no gradient; :func:`rasterize_world_space_gaussians` is the training path for + it. A ``crop`` selects a window of the images: tiles outside it are skipped and the result is + exactly the corresponding region of the uncropped render. The output buffers are still allocated + at full image size before slicing. + + Args: + projected (ProjectedGaussians): Output of :func:`project_gaussians`. + features (torch.Tensor): Per-camera, per-Gaussian features, ``[C, N, D]``. + opacities (torch.Tensor): Per-camera opacities, ``[C, N]``, from :func:`compute_gaussian_opacities`. + Compute them once per render and pass the same tensor to every stage. + tiles (GaussianTileIntersection): Output of :func:`intersect_gaussian_tiles` for ``projected``. + backgrounds (torch.Tensor | None): Per-camera background features, ``[C, D]``. Black if ``None``. + masks (torch.Tensor | None): Boolean per-pixel render mask in image coordinates, ``[C, H, W]``. + Masked-out pixels receive the background with zero alpha and no gradient. + crop (tuple[int, int, int, int] | None): ``(origin_w, origin_h, width, height)`` window to keep, + clipped to the image; a crop entirely outside it yields an empty ``[C, 0, 0, D]`` render, which + :meth:`~fvdb_reality_capture.GaussianSplat3d.render_from_projected_gaussians` pads back to the requested size. + crop_masks (torch.Tensor | None): Boolean per-pixel render mask in crop coordinates, of the crop's + requested ``[C, height, width]`` or clipped size; an alternative to ``masks`` when a crop is given. + + Returns: + images (torch.Tensor): Blended features, ``[C, H, W, D]`` (or the crop size). + alphas (torch.Tensor): Accumulated alpha in ``[0, 1)``, ``[C, H, W, 1]`` (or the crop size). + """ + opacities = check_opacities(opacities, projected) + crop, masks, tile_masks = _render_masks(crop, masks, crop_masks, projected, tiles, opacities.device) + if crop is not None and (crop[2] == 0 or crop[3] == 0): + return _empty_render(features) + images, alphas = cast( + tuple[torch.Tensor, torch.Tensor], + _RasterizeScreenSpaceGaussiansFn.apply( + projected.means2d, + projected.conics, + features, + opacities, + tiles.image_width, + tiles.image_height, + 0, + 0, + tiles.tile_size, + tiles.tile_offsets, + tiles.tile_gaussian_ids, + False, + backgrounds, + tile_masks, + ), + ) + return _finish_dense_render(images, alphas, crop, masks, backgrounds) + + +def rasterize_world_space_gaussians( + means: torch.Tensor, + quats: torch.Tensor, + log_scales: torch.Tensor, + projected: ProjectedGaussians, + features: torch.Tensor, + opacities: torch.Tensor, + world_to_camera_matrices: torch.Tensor, + projection_matrices: torch.Tensor, + tiles: GaussianTileIntersection, + distortion_coeffs: torch.Tensor | None = None, + backgrounds: torch.Tensor | None = None, + masks: torch.Tensor | None = None, + crop: Crop | None = None, + crop_masks: torch.Tensor | None = None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Alpha-blend 3D Gaussians into dense images by evaluating them along per-pixel rays. + + Differentiable with respect to the 3D parameters, ``features`` and ``opacities``, which makes it + the training path for the unscented projection. The projection supplies the tile intersections + and the camera model; the 3D parameters are evaluated directly. + + Args: + means (torch.Tensor): Gaussian centers in world space, ``[N, 3]``. + quats (torch.Tensor): Gaussian rotations as quaternions, ``[N, 4]``. + log_scales (torch.Tensor): Natural-log scale factors, ``[N, 3]``. + projected (ProjectedGaussians): Output of :func:`project_gaussians` for these Gaussians and cameras. + features (torch.Tensor): Per-camera, per-Gaussian features, ``[C, N, D]``. + opacities (torch.Tensor): Per-camera opacities, ``[C, N]``, from :func:`compute_gaussian_opacities`. + world_to_camera_matrices (torch.Tensor): World-to-camera transforms, ``[C, 4, 4]``. + projection_matrices (torch.Tensor): Camera intrinsics, ``[C, 3, 3]``. + tiles (GaussianTileIntersection): Output of :func:`intersect_gaussian_tiles` for ``projected``. + distortion_coeffs (torch.Tensor | None): Packed OpenCV distortion coefficients, ``[C, 12]``. + Required for the OpenCV camera models; ``None`` is allowed for pinhole and orthographic cameras. + backgrounds (torch.Tensor | None): Per-camera background features, ``[C, D]``. Black if ``None``. + masks (torch.Tensor | None): Boolean per-pixel render mask in image coordinates, ``[C, H, W]``. + crop (tuple[int, int, int, int] | None): ``(origin_w, origin_h, width, height)`` window to keep, + clipped to the image; a crop entirely outside it yields an empty ``[C, 0, 0, D]`` render, which + :meth:`~fvdb_reality_capture.GaussianSplat3d.render_from_projected_gaussians` pads back to the requested size. + crop_masks (torch.Tensor | None): Boolean per-pixel render mask in crop coordinates, of the crop's + requested ``[C, height, width]`` or clipped size; an alternative to ``masks`` when a crop is given. + + Returns: + images (torch.Tensor): Blended features, ``[C, H, W, D]`` (or the crop size). + alphas (torch.Tensor): Accumulated alpha in ``[0, 1)``, ``[C, H, W, 1]`` (or the crop size). + """ + opacities = check_opacities(opacities, projected) + distortion_coeffs = check_distortion_coeffs( + distortion_coeffs, projected.camera_model, projected.num_cameras, opacities.device + ) + if distortion_coeffs is None: + distortion_coeffs = torch.zeros( + projected.num_cameras, 12, device=world_to_camera_matrices.device, dtype=world_to_camera_matrices.dtype + ) + crop, masks, tile_masks = _render_masks(crop, masks, crop_masks, projected, tiles, opacities.device) + if crop is not None and (crop[2] == 0 or crop[3] == 0): + return _empty_render(features) + images, alphas = cast( + tuple[torch.Tensor, torch.Tensor], + _RasterizeWorldSpaceGaussiansFn.apply( + means, + quats, + log_scales, + features, + opacities, + world_to_camera_matrices, + world_to_camera_matrices, + projection_matrices, + distortion_coeffs, + int(RollingShutterType.NONE), + int(projected.camera_model), + tiles.image_width, + tiles.image_height, + 0, + 0, + tiles.tile_size, + tiles.tile_offsets, + tiles.tile_gaussian_ids, + backgrounds, + tile_masks, + ), + ) + return _finish_dense_render(images, alphas, crop, masks, backgrounds) + + +def rasterize_screen_space_gaussians_sparse( + projected: ProjectedGaussians, + features: torch.Tensor, + opacities: torch.Tensor, + sparse_tiles: SparseGaussianTileIntersection, + backgrounds: torch.Tensor | None = None, + tile_masks: torch.Tensor | None = None, +) -> tuple[JaggedTensor, JaggedTensor]: + """Alpha-blend projected Gaussians at the requested pixels only. + + Differentiable with respect to ``features``, ``opacities`` and, for an ``ANALYTIC`` projection, the + projection itself. Results are returned in the order of ``sparse_tiles.pixels_to_render``, + duplicates included. + + Args: + projected (ProjectedGaussians): Output of :func:`project_gaussians`. + features (torch.Tensor): Per-camera, per-Gaussian features, ``[C, N, D]``. + opacities (torch.Tensor): Per-camera opacities, ``[C, N]``, from :func:`compute_gaussian_opacities`. + sparse_tiles (SparseGaussianTileIntersection): Output of :func:`intersect_gaussian_tiles_sparse`. + backgrounds (torch.Tensor | None): Per-camera background features, ``[C, D]``. Black if ``None``. + tile_masks (torch.Tensor | None): Boolean per-tile render mask, ``[C, num_tiles_h, num_tiles_w]``. + Unlike the dense rasterizers this is per tile, since the pixels to render are explicit; + use :func:`pixel_mask_to_tile_mask` to derive it from a per-pixel mask. + + Returns: + features (JaggedTensor): Blended features per requested pixel, one ``[P_c, D]`` list per camera. + alphas (JaggedTensor): Accumulated alpha per requested pixel, one ``[P_c, 1]`` list per camera. + """ + opacities = check_opacities(opacities, projected) + check_tiles_match(sparse_tiles, projected) + if tile_masks is not None: + expected = tuple(sparse_tiles.active_tile_mask.shape) + if tuple(tile_masks.shape) != expected: + raise ValueError( + f"tile_masks must be a per-tile [C, tiles_h, tiles_w] mask of shape {expected}, got {tuple(tile_masks.shape)}" + ) + if tile_masks.device != opacities.device: + raise ValueError(f"tile_masks must be on {opacities.device}, got {tile_masks.device}") + rendered, alphas = cast( + tuple[torch.Tensor, torch.Tensor], + _RasterizeScreenSpaceGaussiansSparseFn.apply( + projected.means2d, + projected.conics, + features, + opacities, + sparse_tiles.unique_pixels, + sparse_tiles.image_width, + sparse_tiles.image_height, + 0, + 0, + sparse_tiles.tile_size, + sparse_tiles.tile_offsets, + sparse_tiles.tile_gaussian_ids, + sparse_tiles.active_tiles, + sparse_tiles.tile_pixel_mask, + sparse_tiles.tile_pixel_cumsum, + sparse_tiles.pixel_map, + False, + backgrounds, + None if tile_masks is None else tile_masks.bool(), + ), + ) + requested = sparse_tiles.pixels_to_render + return ( + requested.jagged_like(sparse_tiles.expand_to_requested(rendered)), + requested.jagged_like(sparse_tiles.expand_to_requested(alphas)), + ) diff --git a/fvdb_reality_capture/functional/_spherical_harmonics.py b/fvdb_reality_capture/functional/_spherical_harmonics.py new file mode 100644 index 00000000..42ec0479 --- /dev/null +++ b/fvdb_reality_capture/functional/_spherical_harmonics.py @@ -0,0 +1,104 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +"""Stage 2 of the composable pipeline: view-dependent per-Gaussian features.""" + +from __future__ import annotations + +import math +from typing import cast + +import torch + +from ..enums import GaussianRenderMode +from ._autograd import _EvaluateGaussianSHFn +from ._types import ProjectedGaussians + + +def sh_degree_from_coefficients(shN: torch.Tensor) -> int: + """The spherical-harmonics degree implied by the number of higher-order coefficient bands. + + Args: + shN (torch.Tensor): Higher-order coefficients, ``[N, K - 1, D]`` with ``K = (degree + 1)**2``. + + Returns: + degree (int): The spherical-harmonics degree. + """ + return math.isqrt(shN.shape[1] + 1) - 1 + + +def evaluate_gaussian_sh( + means: torch.Tensor, + sh0: torch.Tensor, + shN: torch.Tensor, + world_to_camera_matrices: torch.Tensor, + projected: ProjectedGaussians, + sh_degree_to_use: int = -1, + render_mode: GaussianRenderMode = GaussianRenderMode.FEATURES, +) -> torch.Tensor: + """Evaluate spherical harmonics into the per-camera, per-Gaussian features to rasterize. + + Gaussians culled by the projection (zero radii) receive zero features. Differentiable with + respect to ``sh0``, ``shN``, ``means`` and ``world_to_camera_matrices``. The depth channel is the + view-space depth of each Gaussian center, computed here from ``means`` and ``world_to_camera_matrices`` + so that the result is the same function of its inputs under either projection and with or without + autograd, and so its backward is an elementwise product rather than the projection's full backward + kernel (the unscented projection has none). It equals ``projected.depths`` for the Gaussians the + projection kept and is zero for culled ones (zero radii), whose projected depth is not defined. + + Args: + means (torch.Tensor): Gaussian centers in world space, ``[N, 3]``. + sh0 (torch.Tensor): Degree-0 coefficients, ``[N, 1, D]``. + shN (torch.Tensor): Higher-order coefficients, ``[N, K - 1, D]``. May have zero bands. + world_to_camera_matrices (torch.Tensor): World-to-camera transforms, ``[C, 4, 4]``. + projected (ProjectedGaussians): Projection of the same Gaussians into the same cameras. + sh_degree_to_use (int): Highest degree to evaluate. ``-1`` uses every band in ``shN``. + render_mode (GaussianRenderMode): Which features to produce; see the enum for shapes. + + Returns: + features (torch.Tensor): ``[C, N, D]``, ``[C, N, 1]`` or ``[C, N, D + 1]`` depending on ``render_mode``. + + Raises: + ValueError: If ``sh_degree_to_use`` exceeds the degree ``shN`` provides. + """ + render_mode = GaussianRenderMode(render_mode) + if render_mode == GaussianRenderMode.FEATURES: + depths = None + else: + # The view-space z of each center, computed here rather than read from the projection so the + # result does not depend on the projection method or on autograd state, and so its backward is + # an elementwise product and a sum rather than the analytic projection's full backward kernel. + rotation_z = world_to_camera_matrices[:, 2, :3] # [C, 3] + translation_z = world_to_camera_matrices[:, 2, 3:4] # [C, 1] + depths = (rotation_z.unsqueeze(1) * means.unsqueeze(0)).sum(-1) + translation_z # [C, N] + # Culled Gaussians get zero depth, as the projection and the feature kernel give them. + visible = (projected.radii > 0).all(-1) + depths = (depths * visible.to(depths.dtype)).unsqueeze(-1) # [C, N, 1] + if render_mode == GaussianRenderMode.DEPTH: + return depths + + available_degree = sh_degree_from_coefficients(shN) + degree = available_degree if sh_degree_to_use < 0 else sh_degree_to_use + if degree > available_degree: + raise ValueError(f"sh_degree_to_use={degree} exceeds the degree {available_degree} available in shN") + if degree == 0: + shN = sh0.new_empty(sh0.shape[0], 0, sh0.shape[2]) + + empty_ids = torch.empty(0, dtype=torch.int32, device=means.device) + features = cast( + torch.Tensor, + _EvaluateGaussianSHFn.apply( + degree, + projected.num_cameras, + means, + world_to_camera_matrices, + empty_ids, + empty_ids, + sh0, + shN, + projected.radii, + ), + ) + if render_mode == GaussianRenderMode.FEATURES_AND_DEPTH: + features = torch.cat([features, depths], dim=-1) + return features diff --git a/fvdb_reality_capture/functional/_tile_intersection.py b/fvdb_reality_capture/functional/_tile_intersection.py new file mode 100644 index 00000000..d6be65f4 --- /dev/null +++ b/fvdb_reality_capture/functional/_tile_intersection.py @@ -0,0 +1,242 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +"""Stage 3 of the composable pipeline: bin projected Gaussians into image tiles.""" + +from __future__ import annotations + +import math + +import torch +from fvdb import JaggedTensor +from fvdb import functional as F +from fvdb.functional import as_pixel_jagged + +from ._opacity import check_opacities +from ._types import GaussianTileIntersection, ProjectedGaussians, SparseGaussianTileIntersection + + +def _tile_grid(image_width: int, image_height: int, tile_size: int) -> tuple[int, int]: + return math.ceil(image_height / tile_size), math.ceil(image_width / tile_size) + + +def check_tiles_match( + tiles: GaussianTileIntersection | SparseGaussianTileIntersection, projected: ProjectedGaussians +) -> None: + """Raise if ``tiles`` were not intersected for ``projected``. + + Tile intersections index the kernels by camera, image tile and Gaussian. Stale ones, from a projection + of a different camera batch or image size, or of the same cameras before an optimizer step or a + refinement moved, added or removed Gaussians, would read out of range or blend the wrong Gaussians + rather than raise, so the stages check this up front. Every projection carries a token that its tile + intersections record; ``dataclasses.replace`` on the projection keeps the token, so a projection + rebuilt with ``replace`` (for example with some fields detached) still matches its tiles. + + Args: + tiles (GaussianTileIntersection | SparseGaussianTileIntersection): The tile intersection to check. + projected (ProjectedGaussians): The projection the tiles are about to be rasterized with. + """ + if (tiles.image_width, tiles.image_height) != (projected.image_width, projected.image_height): + raise ValueError( + f"tiles were intersected for a {tiles.image_width}x{tiles.image_height} image but the projection is " + f"{projected.image_width}x{projected.image_height}" + ) + if isinstance(tiles, GaussianTileIntersection): + num_cameras = tiles.tile_offsets.shape[0] + else: + num_cameras = tiles.pixels_to_render.num_tensors + if num_cameras != projected.num_cameras: + raise ValueError(f"tiles cover {num_cameras} cameras but the projection has {projected.num_cameras}") + if tiles.projection_token is not projected.token: + raise ValueError( + "tiles were intersected for a different projection; recompute them after the Gaussians or cameras change" + ) + + +def _culling_inputs( + projected: ProjectedGaussians, opacities: torch.Tensor | None +) -> tuple[torch.Tensor | None, torch.Tensor | None]: + """Conics and opacities for the tighter iso-contour tile test, or ``None`` for the bounding-box test.""" + if opacities is None: + return None, None + return projected.conics, check_opacities(opacities, projected).detach() + + +def intersect_gaussian_tiles( + projected: ProjectedGaussians, + opacities: torch.Tensor | None = None, + tile_size: int = 16, +) -> GaussianTileIntersection: + """Bin projected Gaussians into the tiles of the full image, sorted by camera, tile and depth. + + Not differentiable. Passing ``opacities`` enables a tighter per-tile culling test based on each + Gaussian's iso-contour at the opacity threshold instead of its bounding box, which reduces the + work of every later stage. + + Args: + projected (ProjectedGaussians): Output of :func:`project_gaussians`. + opacities (torch.Tensor | None): Per-camera opacities, ``[C, N]``, from + :func:`compute_gaussian_opacities`, for tighter culling. ``None`` culls by bounding box. + tile_size (int): Tile side length in pixels. + + Returns: + tiles (GaussianTileIntersection): The tile intersections. + """ + num_tiles_h, num_tiles_w = _tile_grid(projected.image_width, projected.image_height, tile_size) + conics, opacities = _culling_inputs(projected, opacities) + tile_offsets, tile_gaussian_ids = F.intersect_gaussian_tiles( + projected.means2d, + projected.radii, + projected.depths, + projected.num_cameras, + tile_size, + num_tiles_h, + num_tiles_w, + conics=conics, + opacities=opacities, + ) + return GaussianTileIntersection( + tile_offsets=tile_offsets, + tile_gaussian_ids=tile_gaussian_ids, + tile_size=tile_size, + image_width=projected.image_width, + image_height=projected.image_height, + projection_token=projected.token, + ) + + +def deduplicate_pixels( + pixels_to_render: JaggedTensor, image_width: int, image_height: int +) -> tuple[JaggedTensor, torch.Tensor, bool]: + """Remove pixels that appear more than once within a camera. + + Each pixel's first occurrence is kept, the unique pixels are in the order they were first requested, + and pixels outside the image are never merged with anything. + + Args: + pixels_to_render (JaggedTensor): ``(row, col)`` integer pixels, one list per camera. + image_width (int): Image width, used to linearize pixel coordinates. + image_height (int): Image height, used to linearize pixel coordinates. + + Returns: + unique_pixels (JaggedTensor): The pixels with duplicates removed. + inverse_indices (torch.Tensor): For each flat requested pixel, its index into the flat unique pixels. + Empty when there are no duplicates, since nothing needs reordering then. + has_duplicates (bool): Whether anything was removed. When ``False`` the input is returned as is. + """ + jdata = pixels_to_render.jdata + total_pixels = jdata.shape[0] + device = jdata.device + if total_pixels == 0: + return pixels_to_render, torch.empty(0, dtype=torch.long, device=device), False + + jidx = pixels_to_render.jidx + single_list = jidx.shape[0] == 0 + num_lists = pixels_to_render.num_tensors + rows = jdata[:, 0].long() + cols = jdata[:, 1].long() + keys = rows * image_width + cols + if not single_list: + keys = keys + jidx.long() * (image_height * image_width) + # Linearizing (row, col) aliases pixels outside the image onto valid ones, e.g. (0, W) onto (1, 0). + # Give each such pixel a key of its own so it is never merged into a valid pixel's output and the + # layout kernel's bounds check still sees it. + in_image = (rows >= 0) & (rows < image_height) & (cols >= 0) & (cols < image_width) + out_of_image_keys = num_lists * image_height * image_width + torch.arange(total_pixels, device=device) + keys = torch.where(in_image, keys, out_of_image_keys) + + # A stable sort puts each pixel's first occurrence first within its group of duplicates. + sorted_keys, sort_perm = keys.sort(stable=True) + is_group_start = torch.ones(total_pixels, dtype=torch.bool, device=device) + if total_pixels > 1: + is_group_start[1:] = sorted_keys[1:] != sorted_keys[:-1] + group_ids = is_group_start.long().cumsum(0) - 1 + num_unique = int(group_ids[-1].item()) + 1 + if num_unique == total_pixels: + return pixels_to_render, torch.empty(0, dtype=torch.long, device=device), False + + # Order the unique pixels by first request rather than by sorted key, and remap the groups to match. + unique_orig_indices, order = sort_perm[is_group_start].sort() + rank = torch.empty(num_unique, dtype=torch.long, device=device) + rank[order] = torch.arange(num_unique, device=device) + inverse_indices = torch.empty(total_pixels, dtype=torch.long, device=device) + inverse_indices[sort_perm] = rank[group_ids] + unique_jdata = jdata[unique_orig_indices] + + if single_list: + unique_batch_idx = torch.zeros(num_unique, dtype=torch.long, device=device) + else: + unique_batch_idx = jidx.long()[unique_orig_indices] + counts_per_list = torch.bincount(unique_batch_idx, minlength=num_lists) + new_offsets = torch.zeros(num_lists + 1, dtype=torch.long, device=device) + new_offsets[1:] = counts_per_list.cumsum(0) + unique_pixels = JaggedTensor.from_data_and_offsets(unique_jdata, new_offsets) + return unique_pixels, inverse_indices, True + + +def intersect_gaussian_tiles_sparse( + pixels_to_render: JaggedTensor | torch.Tensor, + projected: ProjectedGaussians, + opacities: torch.Tensor | None = None, + tile_size: int = 16, +) -> SparseGaussianTileIntersection: + """Bin projected Gaussians into only the tiles that contain requested pixels. + + Not differentiable. Requested pixels are deduplicated per camera before the layout is built; + the result records how to expand per-pixel outputs back to the requested order. + + Args: + pixels_to_render (JaggedTensor | torch.Tensor): ``(row, col)`` integer pixels, one list per + camera, or a ``[C, P, 2]`` tensor. + projected (ProjectedGaussians): Output of :func:`project_gaussians`. + opacities (torch.Tensor | None): Per-camera opacities, ``[C, N]``, from + :func:`compute_gaussian_opacities`, for tighter culling. ``None`` culls by bounding box. + tile_size (int): Tile side length in pixels. The sparse kernels require ``16``. + + Returns: + sparse_tiles (SparseGaussianTileIntersection): The sparse tile intersections. + """ + pixels = as_pixel_jagged(pixels_to_render) + num_tiles_h, num_tiles_w = _tile_grid(projected.image_width, projected.image_height, tile_size) + unique_pixels, inverse_indices, has_duplicates = deduplicate_pixels( + pixels, projected.image_width, projected.image_height + ) + active_tiles, active_tile_mask, tile_pixel_mask, tile_pixel_cumsum, pixel_map = F.build_sparse_gaussian_tile_layout( + tile_size, + num_tiles_h, + num_tiles_w, + unique_pixels, + image_width=projected.image_width, + image_height=projected.image_height, + ) + conics, opacities = _culling_inputs(projected, opacities) + tile_offsets, tile_gaussian_ids = F.intersect_gaussian_tiles_sparse( + projected.means2d, + projected.radii, + projected.depths, + active_tile_mask, + active_tiles, + projected.num_cameras, + tile_size, + num_tiles_h, + num_tiles_w, + conics=conics, + opacities=opacities, + ) + return SparseGaussianTileIntersection( + tile_offsets=tile_offsets, + tile_gaussian_ids=tile_gaussian_ids, + pixels_to_render=pixels, + unique_pixels=unique_pixels, + inverse_indices=inverse_indices, + has_duplicates=has_duplicates, + active_tiles=active_tiles, + active_tile_mask=active_tile_mask, + tile_pixel_mask=tile_pixel_mask, + tile_pixel_cumsum=tile_pixel_cumsum, + pixel_map=pixel_map, + tile_size=tile_size, + image_width=projected.image_width, + image_height=projected.image_height, + projection_token=projected.token, + ) diff --git a/fvdb_reality_capture/functional/_types.py b/fvdb_reality_capture/functional/_types.py new file mode 100644 index 00000000..ab9067b5 --- /dev/null +++ b/fvdb_reality_capture/functional/_types.py @@ -0,0 +1,183 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +"""Frozen dataclasses passed between the stages of the composable Gaussian splatting pipeline.""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +import torch +from fvdb import JaggedTensor + +from ..enums import CameraModel, ProjectionMethod + + +@dataclass(frozen=True, eq=False) +class ProjectedGaussians: + """ + Output of :func:`project_gaussians`: the 2D footprint of every Gaussian in every camera. + + ``C`` is the number of cameras and ``N`` the number of Gaussians. Gaussians that were culled + (outside the near/far planes or too small) have both radii set to zero; their other entries are + undefined and are ignored by every downstream stage. + """ + + radii: torch.Tensor + """Per-axis projected radii in pixels, ``[C, N, 2]``, ``int32``.""" + + means2d: torch.Tensor + """Projected centers in pixel coordinates, ``[C, N, 2]``.""" + + depths: torch.Tensor + """View-space depths, ``[C, N]``.""" + + conics: torch.Tensor + """Inverse 2D covariances as ``(a, b, c)`` of ``a x^2 + 2 b x y + c y^2``, ``[C, N, 3]``.""" + + compensations: torch.Tensor | None + """Anti-aliasing opacity compensation factors, ``[C, N]``, or ``None`` when ``antialias`` was off.""" + + image_width: int + """Width in pixels of the images the Gaussians were projected into.""" + + image_height: int + """Height in pixels of the images the Gaussians were projected into.""" + + camera_model: CameraModel + """Camera model used for the projection.""" + + projection_method: ProjectionMethod + """The projection method actually used, with :attr:`~fvdb_reality_capture.ProjectionMethod.AUTO` resolved.""" + + token: object = field(default_factory=object, repr=False) + """Identity of this projection. Tile intersections record it so stale tiles are caught; ``replace`` keeps it.""" + + @property + def num_cameras(self) -> int: + """Number of cameras ``C``.""" + return self.means2d.shape[0] + + @property + def num_gaussians(self) -> int: + """Number of Gaussians ``N``.""" + return self.means2d.shape[1] + + @property + def is_differentiable(self) -> bool: + """Whether gradients flow from the 2D quantities back to the 3D Gaussian parameters. + + This is ``True`` for the analytic projection and ``False`` for the unscented transform, whose + kernel has no backward pass. World-space rasterization differentiates through the 3D + parameters directly and does not need it. + """ + return self.projection_method == ProjectionMethod.ANALYTIC + + +@dataclass(frozen=True, eq=False) +class GaussianTileIntersection: + """Output of :func:`intersect_gaussian_tiles`: which Gaussians touch which image tile, depth sorted.""" + + tile_offsets: torch.Tensor + """Start index into :attr:`tile_gaussian_ids` for each tile, ``[C, num_tiles_h, num_tiles_w]``.""" + + tile_gaussian_ids: torch.Tensor + """Flattened Gaussian index of every tile intersection, ``[num_intersections]``.""" + + tile_size: int + """Tile side length in pixels.""" + + image_width: int + """Image width in pixels.""" + + image_height: int + """Image height in pixels.""" + + projection_token: object = field(repr=False) + """The :attr:`ProjectedGaussians.token` of the projection these tiles were intersected for.""" + + @property + def num_tiles_w(self) -> int: + """Number of tiles along the image width.""" + return self.tile_offsets.shape[2] + + @property + def num_tiles_h(self) -> int: + """Number of tiles along the image height.""" + return self.tile_offsets.shape[1] + + +@dataclass(frozen=True, eq=False) +class SparseGaussianTileIntersection: + """ + Output of :func:`intersect_gaussian_tiles_sparse`: tile bookkeeping for rendering an arbitrary + set of pixels. + + The sparse kernels require each pixel to appear once per camera, so the requested pixels are + deduplicated here. Downstream stages render :attr:`unique_pixels` and use + :attr:`inverse_indices` to expand their results back to the order of :attr:`pixels_to_render`. + """ + + tile_offsets: torch.Tensor + """Start index into :attr:`tile_gaussian_ids` per active tile plus a trailing end, ``[AT + 1]``.""" + + tile_gaussian_ids: torch.Tensor + """Flattened Gaussian index of every tile intersection, ``[num_intersections]``.""" + + pixels_to_render: JaggedTensor + """The requested ``(row, col)`` pixels, one list per camera, normalized to a JaggedTensor by :func:`as_pixel_jagged`.""" + + unique_pixels: JaggedTensor + """The requested pixels with per-camera duplicates removed. Equal to :attr:`pixels_to_render` if none.""" + + inverse_indices: torch.Tensor + """Index into the flat unique pixels for each flat requested pixel, ``[num_requested]``. Empty when + :attr:`has_duplicates` is ``False``, since nothing needs reordering then.""" + + has_duplicates: bool + """Whether :attr:`pixels_to_render` contained duplicate pixels within a camera.""" + + active_tiles: torch.Tensor + """Flattened ids of the tiles that contain at least one requested pixel, ``[AT]``.""" + + active_tile_mask: torch.Tensor + """Boolean mask of active tiles, ``[C, num_tiles_h, num_tiles_w]``.""" + + tile_pixel_mask: torch.Tensor + """Per active tile bitmask of the requested pixels within it, ``[AT, words_per_tile]``, ``uint64``.""" + + tile_pixel_cumsum: torch.Tensor + """Inclusive cumulative count of requested pixels over the active tiles, ``[AT]``.""" + + pixel_map: torch.Tensor + """Output slot of the ``k``-th requested pixel of each active tile, ``[num_unique]``.""" + + tile_size: int + """Tile side length in pixels.""" + + image_width: int + """Image width in pixels.""" + + image_height: int + """Image height in pixels.""" + + projection_token: object = field(repr=False) + """The :attr:`ProjectedGaussians.token` of the projection these tiles were intersected for.""" + + @property + def num_cameras(self) -> int: + """Number of cameras ``C``.""" + return self.active_tile_mask.shape[0] + + def expand_to_requested(self, per_unique_pixel: torch.Tensor) -> torch.Tensor: + """Reorder a flat per-unique-pixel result into the order of :attr:`pixels_to_render`. + + Args: + per_unique_pixel (torch.Tensor): Tensor whose first dimension runs over the unique pixels. + + Returns: + per_requested_pixel (torch.Tensor): The same values indexed by the requested pixels. + """ + if not self.has_duplicates: + return per_unique_pixel + return per_unique_pixel.index_select(0, self.inverse_indices) diff --git a/fvdb_reality_capture/radiance_fields/gaussian_splatting.py b/fvdb_reality_capture/radiance_fields/gaussian_splatting.py index 208308c9..83cf8160 100644 --- a/fvdb_reality_capture/radiance_fields/gaussian_splatting.py +++ b/fvdb_reality_capture/radiance_fields/gaussian_splatting.py @@ -5,72 +5,46 @@ import math import pathlib -from typing import Any, Mapping, Sequence, TypeVar, cast, overload +from typing import Any, Mapping, Sequence, TypeVar, overload import torch -import torch.nn.functional as F - -from fvdb import _fvdb_cpp as _C -from fvdb._fvdb_cpp import JaggedTensor as JaggedTensorCpp -from ._gaussian_autograd import ( - _EvaluateGaussianSHFn, - _ProjectGaussiansJaggedFn, - _ProjectGaussiansFn, - _RasterizeScreenSpaceGaussiansFn, - _RasterizeScreenSpaceGaussiansSparseFn, - _RasterizeWorldSpaceGaussiansFn, -) +from fvdb import functional as fvdb_functional from fvdb.grid import Grid from fvdb.grid_batch import GridBatch from fvdb.jagged_tensor import JaggedTensor from fvdb.types import DeviceIdentifier, cast_check, resolve_device -from ..enums import CameraModel, ProjectionMethod +from ..enums import CameraModel, GaussianRenderMode, ProjectionMethod +from ..functional import ( + Crop, + GaussianTileIntersection, + ProjectedGaussians, + as_pixel_jagged, + compute_gaussian_opacities, + evaluate_gaussian_sh, + intersect_gaussian_tiles, + intersect_gaussian_tiles_sparse, + project_gaussians, + rasterize_contributing_gaussian_ids, + rasterize_contributing_gaussian_ids_sparse, + rasterize_num_contributing_gaussians, + rasterize_num_contributing_gaussians_sparse, + rasterize_screen_space_gaussians, + rasterize_screen_space_gaussians_sparse, + pad_crop, + rasterize_world_space_gaussians, + sh_degree_from_coefficients, + validate_crop, +) +from ..functional._autograd import ( + _EvaluateGaussianSHFn, + _ProjectGaussiansJaggedFn, + _RasterizeScreenSpaceGaussiansFn, +) JaggedTensorOrTensorT = TypeVar("JaggedTensorOrTensorT", JaggedTensor, torch.Tensor) -def _pixel_mask_to_tile_mask(pixel_mask: torch.Tensor, tile_size: int) -> torch.Tensor: - """Convert a per-pixel boolean mask ``[C, H, W]`` to a per-tile boolean mask ``[C, tileH, tileW]``. - - A tile is ``True`` (render) if **any** pixel in that tile is ``True``. - Uses ``max_pool2d`` with ``ceil_mode=True`` so that partial edge tiles are - handled correctly when ``H`` or ``W`` is not divisible by ``tile_size``. - """ - return ( - F.max_pool2d( - pixel_mask.unsqueeze(1).float(), - kernel_size=tile_size, - stride=tile_size, - ceil_mode=True, - ) - .bool() - .squeeze(1) - ) - - -def _apply_pixel_mask( - features: torch.Tensor, - alphas: torch.Tensor, - pixel_mask: torch.Tensor, - backgrounds: torch.Tensor | None, -) -> tuple[torch.Tensor, torch.Tensor]: - """Apply a per-pixel boolean mask ``[C, H, W]`` to rendered features and alphas. - - Masked-out pixels (``False``) are filled with the background colour (or zero) - and their alpha is set to zero. The operation is differentiable: gradients - flow through unmasked pixels and are zero for masked pixels. - """ - mask_float = pixel_mask.unsqueeze(-1).float() # [C, H, W, 1] - if backgrounds is not None: - bg = backgrounds[:, None, None, :] # [C, 1, 1, D] - else: - bg = torch.zeros(1, 1, 1, features.shape[-1], device=features.device, dtype=features.dtype) - features = features * mask_float + bg * (1.0 - mask_float) - alphas = alphas * mask_float - return features, alphas - - class ProjectedGaussianSplats: """ A class representing a set of Gaussian splats projected onto a batch of 2D image planes. @@ -90,49 +64,76 @@ class ProjectedGaussianSplats: def __init__( self, *, - radii: torch.Tensor, - means2d: torch.Tensor, - depths: torch.Tensor, - conics: torch.Tensor, - compensations: torch.Tensor | None, + projected: ProjectedGaussians, render_quantities: torch.Tensor, - opacities: torch.Tensor, - image_width: int, - image_height: int, + logit_opacities: torch.Tensor, antialias: bool, eps_2d: float, near_plane: float, far_plane: float, min_radius_2d: float, sh_degree_to_use: int, - camera_model: CameraModel, - projection_method: ProjectionMethod, + opacities: torch.Tensor, _private: Any = None, ) -> None: """ Private constructor. Use :meth:`GaussianSplat3d.project_gaussians_for_images` or similar methods to create instances. + + ``opacities`` are the per-camera opacities of ``projected``, ``(C, N)``, computed by the projection + helper alongside the features so the projection is a consistent snapshot of the model. """ if _private is not self.__PRIVATE__: raise ValueError( "ProjectedGaussianSplats constructor is private. Use GaussianSplat3d.project_gaussians_for_images or similar methods instead." ) - self._radii = radii - self._means2d = means2d - self._depths = depths - self._conics = conics - self._compensations = compensations + self._projected = projected self._render_quantities = render_quantities - self._opacities = opacities - self._image_width = image_width - self._image_height = image_height + self._logit_opacities = logit_opacities self._antialias = antialias self._eps_2d = eps_2d self._near_plane = near_plane self._far_plane = far_plane self._min_radius_2d = min_radius_2d self._sh_degree_to_use = sh_degree_to_use - self._camera_model = camera_model - self._projection_method = projection_method + self._opacities = opacities + + @property + def projected_gaussians(self) -> ProjectedGaussians: + """ + Return the underlying :class:`fvdb_reality_capture.functional.ProjectedGaussians`, the stage-1 output of the + composable pipeline, for use with the functions in :mod:`fvdb_reality_capture.functional`. + + Returns: + projected_gaussians (ProjectedGaussians): The projected Gaussians without features or opacities. + """ + return self._projected + + def tile_intersection(self, tile_size: int = 16) -> GaussianTileIntersection: + """ + Compute the tile intersections of the projected Gaussians for a tile size. + + Nothing is cached, so holding a projection does not hold tile buffers. A caller rendering several + crops from one projection should keep the result and pass it to :meth:`GaussianSplat3d.render_from_projected_gaussians` + as ``tiles`` (see its docstring example) or to the rasterization stages in :mod:`fvdb_reality_capture.functional`. + + Args: + tile_size (int): The tile side length in pixels. Default is 16. + + Returns: + tiles (GaussianTileIntersection): The tile intersections, as consumed by the rasterization stages + in :mod:`fvdb_reality_capture.functional`. + """ + return intersect_gaussian_tiles(self._projected, tile_size=tile_size, opacities=self.opacities) + + @property + def logit_opacities(self) -> torch.Tensor: + """ + Return the logit opacities of the Gaussians that were projected. + + Returns: + logit_opacities (torch.Tensor): A tensor of shape ``(N,)`` where ``N`` is the number of projected Gaussians. + """ + return self._logit_opacities @property def antialias(self) -> bool: @@ -153,11 +154,11 @@ def inv_covar_2d(self) -> torch.Tensor: where each covariance matrix is represented as ``(Cxx, Cxy, Cyy)``. Returns: - inv_covar_2d (torch.Tensor): A tensor of shape ``(C, N, D)`` representing the packed inverse 2D covariance matrices, - where ``C`` is the number of image planes, ``N`` is the number of projected Gaussians, and ``D`` is number of feature channels for each - Gaussian (see :attr:`GaussianSplat3d.num_channels`). + inv_covar_2d (torch.Tensor): A tensor of shape ``(C, N, 3)`` representing the packed inverse 2D covariance matrices, + where ``C`` is the number of image planes, ``N`` is the number of projected Gaussians, and the last dimension holds + ``(a, b, c)`` of the inverse covariance ``a x^2 + 2 b x y + c y^2``. """ - return self._conics + return self._projected.conics @property def depths(self) -> torch.Tensor: @@ -169,7 +170,7 @@ def depths(self) -> torch.Tensor: depths (torch.Tensor): A tensor of shape ``(C, N)`` representing the depth of each projected Gaussian, where ``C`` is the number of image planes, and ``N`` is the number of projected Gaussians. """ - return self._depths + return self._projected.depths @property def eps_2d(self) -> float: @@ -200,7 +201,7 @@ def image_height(self) -> int: Returns: image_height (int): The height of the image planes. """ - return self._image_height + return self._projected.image_height @property def image_width(self) -> int: @@ -210,7 +211,7 @@ def image_width(self) -> int: Returns: image_width (int): The width of the image planes. """ - return self._image_width + return self._projected.image_width @property def means2d(self) -> torch.Tensor: @@ -222,7 +223,7 @@ def means2d(self) -> torch.Tensor: where ``C`` is the number of image planes, ``N`` is the number of projected Gaussians, and the last dimension contains the (x, y) coordinates of the means in pixel space. """ - return self._means2d + return self._projected.means2d @property def min_radius_2d(self) -> float: @@ -250,6 +251,9 @@ def opacities(self) -> torch.Tensor: """ Return the opacities of each projected Gaussian in each image plane. + They are computed when the Gaussians are projected, together with the features, so they reflect + the model at that moment and carry a gradient exactly when the projection was made with one. + Returns: opacities (torch.Tensor): A tensor of shape ``(C, N)`` representing the opacity of each projected Gaussian, where ``C`` is the number of image planes, and ``N`` is the number of projected Gaussians. @@ -264,7 +268,7 @@ def camera_model(self) -> CameraModel: Returns: camera_model (CameraModel): The camera model used during projection. """ - return self._camera_model + return self._projected.camera_model @property def projection_method(self) -> ProjectionMethod: @@ -274,7 +278,7 @@ def projection_method(self) -> ProjectionMethod: Returns: projection_method (ProjectionMethod): The resolved projection method. """ - return self._projection_method + return self._projected.projection_method @property def radii(self) -> torch.Tensor: @@ -287,7 +291,7 @@ def radii(self) -> torch.Tensor: radii (torch.Tensor): A tensor of shape ``(C, N, 2)`` representing the per-axis 2D radius of each projected Gaussian. """ - return self._radii + return self._projected.radii @property def render_quantities(self) -> torch.Tensor: @@ -534,11 +538,9 @@ def from_ply( device = resolve_device(device) if isinstance(filename, pathlib.Path): filename = str(filename) - - means, quats, log_scales, logit_opacities, sh0, shN, metadata = _C.load_gaussian_ply( - filename=filename, device=device + means, quats, log_scales, logit_opacities, sh0, shN, metadata = fvdb_functional.load_gaussian_ply( + filename, device ) - return ( cls( means=means, @@ -937,7 +939,7 @@ def sh_degree(self) -> int: Returns: sh_degree (int): The degree of the spherical harmonics. """ - return int(math.isqrt(self._shN.size(1) + 1)) - 1 + return sh_degree_from_coefficients(self._shN) @property def num_channels(self) -> int: @@ -1325,6 +1327,8 @@ def accumulated_max_2d_radii(self) -> torch.Tensor: If :this :class:`GaussianSplat3d` instance is set to track maximum 2D radii (*i.e* :attr:`accumulate_max_2d_radii` is ``True``), then this tensor contains the maximum 2D radius for each Gaussian. + The projection kernel records radii only alongside the 2D mean-gradient statistics, so this is updated + only when :attr:`accumulate_mean_2d_gradients` is also ``True`` and a backward pass reaches the projected means. If :attr:`accumulate_max_2d_radii` is ``False``, this property will be an empty tensor. @@ -1445,493 +1449,238 @@ def accumulated_mean_2d_gradient_norms(self) -> torch.Tensor: # Private rendering helpers # --------------------------------------------------------------------------- - @staticmethod - def _is_ortho(camera_model: CameraModel) -> bool: - return camera_model == CameraModel.ORTHOGRAPHIC - - @staticmethod - def _resolve_projection_method(camera_model: CameraModel, projection_method: ProjectionMethod) -> ProjectionMethod: - if projection_method != ProjectionMethod.AUTO: - return projection_method - if camera_model in (CameraModel.PINHOLE, CameraModel.ORTHOGRAPHIC): - return ProjectionMethod.ANALYTIC - return ProjectionMethod.UNSCENTED - - @staticmethod - def _use_ut(camera_model: CameraModel, projection_method: ProjectionMethod) -> bool: - return GaussianSplat3d._resolve_projection_method(camera_model, projection_method) == ProjectionMethod.UNSCENTED - - def _do_projection( - self, - w2c: torch.Tensor, - K: torch.Tensor, - W: int, - H: int, - eps2d: float, - near: float, - far: float, - min_radius: float, - antialias: bool, - camera_model: CameraModel, - projection_method: ProjectionMethod, - distortion_coeffs: torch.Tensor | None, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]: - """Project Gaussians onto image planes. - - Returns ``(radii, means2d, depths, conics, compensations)``. - """ - means = self._means - quats = self._quats - log_scales = self._log_scales - ortho = self._is_ortho(camera_model) - - C = w2c.size(0) - if not K.is_contiguous(): - raise RuntimeError("projectionMatrices must be contiguous") - if not w2c.is_contiguous(): - raise RuntimeError("worldToCameraMatrices must be contiguous") - if distortion_coeffs is not None: - if list(distortion_coeffs.shape) != [C, 12]: - raise RuntimeError(f"distortionCoeffs must have shape ({C}, 12)") - if not distortion_coeffs.is_contiguous(): - raise RuntimeError("distortionCoeffs must be contiguous") - - is_opencv = camera_model not in (CameraModel.PINHOLE, CameraModel.ORTHOGRAPHIC) - if is_opencv: - resolved = self._resolve_projection_method(camera_model, projection_method) - if resolved != ProjectionMethod.UNSCENTED: - raise RuntimeError("OpenCV camera models require ProjectionMethod::UNSCENTED or AUTO") - if distortion_coeffs is None: - raise RuntimeError("distortionCoeffs must be provided for OpenCV camera models") - - N = means.size(0) - accum_grad_norms: torch.Tensor | None = None - accum_step_counts: torch.Tensor | None = None - accum_max_radii: torch.Tensor | None = None - + def _projection_accumulators(self) -> tuple[torch.Tensor | None, torch.Tensor | None, torch.Tensor | None]: + """Create the enabled densification accumulators if missing or stale and return them.""" + num_gaussians = self.num_gaussians + device = self._means.device + grad_norms: torch.Tensor | None = None + step_counts: torch.Tensor | None = None + max_radii: torch.Tensor | None = None if self._accumulate_mean_2d_gradients: gn = self._accumulated_mean_2d_gradient_norms - if gn is None or gn.numel() != N: - gn = torch.zeros(N, device=means.device, dtype=means.dtype) + if gn is None or gn.numel() != num_gaussians: + gn = torch.zeros(num_gaussians, device=device, dtype=self._means.dtype) self._accumulated_mean_2d_gradient_norms = gn - accum_grad_norms = gn - sc = self._accumulated_gradient_step_counts - if sc is None or sc.numel() != N: - sc = torch.zeros(N, device=means.device, dtype=torch.int32) + if sc is None or sc.numel() != num_gaussians: + sc = torch.zeros(num_gaussians, device=device, dtype=torch.int32) self._accumulated_gradient_step_counts = sc - accum_step_counts = sc - + grad_norms, step_counts = gn, sc if self._accumulate_max_2d_radii: mr = self._accumulated_max_2d_radii - if mr is None or mr.numel() != N: - mr = torch.zeros(N, device=means.device, dtype=torch.int32) + if mr is None or mr.numel() != num_gaussians: + mr = torch.zeros(num_gaussians, device=device, dtype=torch.int32) self._accumulated_max_2d_radii = mr - accum_max_radii = mr - - if self._use_ut(camera_model, projection_method): - if distortion_coeffs is None: - distortion_coeffs = torch.empty(C, 0, device=means.device, dtype=means.dtype) - result = _C.project_gaussians_ut_fwd( - means, - quats, - log_scales, - w2c, - w2c, - K, - distortion_coeffs, - self._camera_model_to_cpp(camera_model), - W, - H, - eps2d, - near, - far, - min_radius, - antialias, - ) - radii, means2d, depths, conics, compensations = result - if not antialias: - compensations = None - return radii, means2d, depths, conics, compensations - - result = _ProjectGaussiansFn.apply( - means, - quats, - log_scales, - w2c, - K, - W, - H, - eps2d, - near, - far, - min_radius, - antialias, - ortho, - accum_grad_norms, - accum_step_counts, - accum_max_radii, - ) - radii = result[0] - means2d = result[1] - depths = result[2] - conics = result[3] - compensations = result[4] if antialias and len(result) > 4 else None - return radii, means2d, depths, conics, compensations - - def _eval_sh( - self, - w2c: torch.Tensor, - radii: torch.Tensor, - sh_degree_to_use: int, - ) -> torch.Tensor: - """Evaluate spherical harmonics to produce per-Gaussian color features ``[C, N, D]``.""" - means = self._means - sh0 = self._sh0 - shN = self._shN - C = w2c.size(0) - - sh_degree = self.sh_degree - if sh_degree_to_use < 0: - sh_degree_to_use = sh_degree - - if sh_degree_to_use > 0: - empty_ids = torch.empty(0, dtype=torch.int32, device=means.device) - return _EvaluateGaussianSHFn.apply( - sh_degree_to_use, - C, - means, - w2c, - empty_ids, - empty_ids, - sh0, - shN, - radii, - ) - else: - shN = sh0.new_empty(sh0.shape[0], 0, sh0.shape[2]) - empty_ids = torch.empty(0, dtype=torch.int32, device=means.device) - return _EvaluateGaussianSHFn.apply( - sh_degree_to_use, - C, - means, - w2c, - empty_ids, - empty_ids, - sh0, - shN, - radii, - ) - - def _make_render_features( - self, - w2c: torch.Tensor, - radii: torch.Tensor, - depths: torch.Tensor, - sh_degree_to_use: int, - include_colors: bool, - include_depth: bool, - ) -> torch.Tensor: - """Build the feature tensor used for rasterization. + max_radii = mr + return grad_norms, step_counts, max_radii - ``include_colors=True, include_depth=False`` -> ``[C, N, D]`` (colors) - ``include_colors=False, include_depth=True`` -> ``[C, N, 1]`` (depth) - ``include_colors=True, include_depth=True`` -> ``[C, N, D+1]`` (colors + depth) - """ - parts: list[torch.Tensor] = [] - if include_colors: - parts.append(self._eval_sh(w2c, radii, sh_degree_to_use)) - if include_depth: - parts.append(depths.unsqueeze(-1)) - return torch.cat(parts, dim=-1) if len(parts) > 1 else parts[0] - - def _make_opacities( + def _project( self, - C: int, - compensations: torch.Tensor | None, + world_to_camera_matrices: torch.Tensor, + projection_matrices: torch.Tensor, + image_width: int, + image_height: int, + near: float, + far: float, + camera_model: CameraModel, + projection_method: ProjectionMethod, + distortion_coeffs: torch.Tensor | None, + min_radius_2d: float, + eps_2d: float, antialias: bool, - ) -> torch.Tensor: - """Sigmoid of logit_opacities, optionally scaled by antialias compensations.""" - - # Ideally, we would like to avoid materializing the repeated [C,N] tensor when opacities - # are shared across cameras by replacing the repeat call with .unsqueeze(0).expand(C, -1). - # However, a non-contiguous opacities tensor is not currently supported in world space - # rasterization and mGPU image space rasterization. - opacities = torch.sigmoid(self._logit_opacities).repeat(C, 1) - if antialias and compensations is not None: - opacities = opacities * compensations - return opacities - - def _intersect_tiles( - self, - means2d: torch.Tensor, - radii: torch.Tensor, - depths: torch.Tensor, - conics: torch.Tensor, - opacities: torch.Tensor, - C: int, - tile_size: int, - W: int, - H: int, - ) -> tuple[torch.Tensor, torch.Tensor, int, int]: - """Compute tile-Gaussian intersections. - - Returns ``(tile_offsets, tile_gaussian_ids, num_tiles_h, num_tiles_w)``. - """ - num_tiles_h = math.ceil(H / tile_size) - num_tiles_w = math.ceil(W / tile_size) - tile_offsets, tile_gaussian_ids = _C.intersect_gaussian_tiles( - means2d, - radii, - depths, - C, - tile_size, - num_tiles_h, - num_tiles_w, - conics=conics, - opacities=opacities, + ) -> ProjectedGaussians: + """Stage 1 for this model's Gaussians, wiring in the enabled densification accumulators.""" + # Every projection wires in the enabled accumulators. World-space rendering reaches the analytic + # backward only through the antialiasing compensations, with a zero 2D-mean gradient, and the + # kernel still counts that as a step. + grad_norms, step_counts, max_radii = self._projection_accumulators() + return project_gaussians( + self._means, + self._quats, + self._log_scales, + world_to_camera_matrices, + projection_matrices, + image_width, + image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + accumulated_mean_2d_gradient_norms=grad_norms, + accumulated_gradient_step_counts=step_counts, + accumulated_max_2d_radii=max_radii, ) - return tile_offsets, tile_gaussian_ids, num_tiles_h, num_tiles_w - def _intersect_tiles_sparse( + def _project_and_opacities( self, - pixels_jt: JaggedTensor, - means2d: torch.Tensor, - radii: torch.Tensor, - depths: torch.Tensor, - conics: torch.Tensor, - opacities: torch.Tensor, - C: int, - tile_size: int, - W: int, - H: int, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - """Compute sparse tile-Gaussian intersections for a set of pixel coordinates. - - Returns ``(tile_offsets, tile_gaussian_ids, active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map)``. - """ - 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, - ) - tile_offsets, tile_gaussian_ids = _C.intersect_gaussian_tiles_sparse( - means2d, - radii, - depths, - active_tile_mask, - active_tiles, - C, - tile_size, - num_tiles_h, - num_tiles_w, - conics=conics, - opacities=opacities, + world_to_camera_matrices: torch.Tensor, + projection_matrices: torch.Tensor, + image_width: int, + image_height: int, + near: float, + far: float, + camera_model: CameraModel, + projection_method: ProjectionMethod, + distortion_coeffs: torch.Tensor | None, + min_radius_2d: float, + eps_2d: float, + antialias: bool, + ) -> tuple[ProjectedGaussians, torch.Tensor]: + """Stage 1 plus the per-camera opacities every later stage takes, computed once.""" + projected = self._project( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, ) - return tile_offsets, tile_gaussian_ids, active_tiles, tile_pixel_mask, tile_pixel_cumsum, pixel_map + return projected, compute_gaussian_opacities(self._logit_opacities, projected) - def _rasterize_screen_space( + def _features( self, - means2d: torch.Tensor, - conics: torch.Tensor, - features: torch.Tensor, - opacities: torch.Tensor, - W: int, - H: int, - tile_size: int, - tile_offsets: torch.Tensor, - tile_gaussian_ids: torch.Tensor, - backgrounds: torch.Tensor | None, - tile_masks: torch.Tensor | None, - ) -> tuple[torch.Tensor, torch.Tensor]: - return cast( - tuple[torch.Tensor, torch.Tensor], - _RasterizeScreenSpaceGaussiansFn.apply( - means2d, - conics, - features, - opacities, - W, - H, - 0, - 0, - tile_size, - tile_offsets, - tile_gaussian_ids, - False, - backgrounds, - tile_masks, - ), + projected: ProjectedGaussians, + world_to_camera_matrices: torch.Tensor, + sh_degree_to_use: int, + render_mode: GaussianRenderMode, + ) -> torch.Tensor: + """Stage 2 for this model's spherical-harmonics coefficients.""" + return evaluate_gaussian_sh( + self._means, self._sh0, self._shN, world_to_camera_matrices, projected, sh_degree_to_use, render_mode ) - def _rasterize_screen_space_sparse( + def _project_for( self, - pixels_jt: JaggedTensor, - means2d: torch.Tensor, - conics: torch.Tensor, - features: torch.Tensor, - opacities: torch.Tensor, - W: int, - H: int, - tile_size: int, - tile_offsets: torch.Tensor, - tile_gaussian_ids: torch.Tensor, - active_tiles: torch.Tensor, - tile_pixel_mask: torch.Tensor, - tile_pixel_cumsum: torch.Tensor, - pixel_map: torch.Tensor, - backgrounds: torch.Tensor | None, - masks: torch.Tensor | None, - ) -> tuple[torch.Tensor, torch.Tensor]: - return cast( - tuple[torch.Tensor, torch.Tensor], - _RasterizeScreenSpaceGaussiansSparseFn.apply( - means2d, - conics, - features, - opacities, - pixels_jt, - W, - H, - 0, - 0, - tile_size, - tile_offsets, - tile_gaussian_ids, - active_tiles, - tile_pixel_mask, - tile_pixel_cumsum, - pixel_map, - False, - backgrounds, - masks, - ), + world_to_camera_matrices: torch.Tensor, + projection_matrices: torch.Tensor, + image_width: int, + image_height: int, + near: float, + far: float, + camera_model: CameraModel, + projection_method: ProjectionMethod, + distortion_coeffs: torch.Tensor | None, + min_radius_2d: float, + eps_2d: float, + antialias: bool, + sh_degree_to_use: int, + render_mode: GaussianRenderMode, + ) -> ProjectedGaussianSplats: + """Stages 1 and 2, bundled for later rendering with :meth:`render_from_projected_gaussians`.""" + projected, opacities = self._project_and_opacities( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + ) + render_quantities = self._features(projected, world_to_camera_matrices, sh_degree_to_use, render_mode) + return ProjectedGaussianSplats( + projected=projected, + render_quantities=render_quantities, + logit_opacities=self._logit_opacities, + antialias=antialias, + eps_2d=eps_2d, + near_plane=near, + far_plane=far, + min_radius_2d=min_radius_2d, + sh_degree_to_use=sh_degree_to_use, + opacities=opacities, + _private=ProjectedGaussianSplats.__PRIVATE__, ) - def _rasterize_world_space( + def _render_dense( self, - features: torch.Tensor, - opacities: torch.Tensor, - w2c: torch.Tensor, - K: torch.Tensor, - distortion_coeffs: torch.Tensor, + world_to_camera_matrices: torch.Tensor, + projection_matrices: torch.Tensor, + image_width: int, + image_height: int, + near: float, + far: float, camera_model: CameraModel, - W: int, - H: int, + projection_method: ProjectionMethod, + distortion_coeffs: torch.Tensor | None, + sh_degree_to_use: int, tile_size: int, - tile_offsets: torch.Tensor, - tile_gaussian_ids: torch.Tensor, + min_radius_2d: float, + eps_2d: float, + antialias: bool, backgrounds: torch.Tensor | None, - tile_masks: torch.Tensor | None, + masks: torch.Tensor | None, + render_mode: GaussianRenderMode, + world_space: bool, + crop: Crop | None = None, + crop_masks: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: - return cast( - tuple[torch.Tensor, torch.Tensor], - _RasterizeWorldSpaceGaussiansFn.apply( + """All four stages for dense images, in screen space or world space.""" + projected, opacities = self._project_and_opacities( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + ) + features = self._features(projected, world_to_camera_matrices, sh_degree_to_use, render_mode) + tiles = intersect_gaussian_tiles(projected, tile_size=tile_size, opacities=opacities) + if world_space: + return rasterize_world_space_gaussians( self._means, self._quats, self._log_scales, + projected, features, opacities, - w2c, - w2c, - K, - distortion_coeffs, - _C.RollingShutterType.NONE.value, - self._camera_model_to_cpp(camera_model).value, - W, - H, - 0, - 0, - tile_size, - tile_offsets, - tile_gaussian_ids, - backgrounds, - tile_masks, - ), + world_to_camera_matrices, + projection_matrices, + tiles, + distortion_coeffs=distortion_coeffs, + backgrounds=backgrounds, + masks=masks, + crop=crop, + crop_masks=crop_masks, + ) + return rasterize_screen_space_gaussians( + projected, + features, + opacities, + tiles, + backgrounds=backgrounds, + masks=masks, + crop=crop, + crop_masks=crop_masks, ) - @staticmethod - def _deduplicate_pixels( - pixels_jt: JaggedTensor, - image_width: int, - image_height: int, - ) -> tuple[JaggedTensor, torch.Tensor, bool]: - """Deduplicate pixel coordinates in a JaggedTensor. - - Returns ``(unique_pixels, inverse_indices, has_duplicates)``. - """ - jdata = pixels_jt.jdata - total_pixels = jdata.shape[0] - - if total_pixels == 0: - empty_inverse = torch.empty(0, dtype=torch.long, device=jdata.device) - return pixels_jt, empty_inverse, False - - device = jdata.device - jidx = pixels_jt.jidx - num_pixels_per_image = image_height * image_width - - single_list = jidx.shape[0] == 0 - if jdata.dtype == torch.int32: - rows = jdata[:, 0].to(torch.long) - cols = jdata[:, 1].to(torch.long) - else: - rows = jdata[:, 0] - cols = jdata[:, 1] - - if single_list: - keys = rows * image_width + cols - else: - keys = jidx.to(torch.long) * num_pixels_per_image + rows * image_width + cols - - sorted_keys, sort_perm = keys.sort() - - is_group_start = torch.ones(total_pixels, dtype=torch.bool, device=device) - if total_pixels > 1: - is_group_start[1:] = sorted_keys[1:] != sorted_keys[:-1] - - first_in_sorted = is_group_start.nonzero(as_tuple=False).squeeze(1) - - group_ids = is_group_start.to(torch.long).cumsum_(0).sub_(1) - num_unique = int(group_ids[-1].item()) + 1 - - if num_unique == total_pixels: - return pixels_jt, torch.arange(total_pixels, dtype=torch.long, device=device), False - - inverse_indices = torch.empty(total_pixels, dtype=torch.long, device=device) - inverse_indices[sort_perm] = group_ids - - unique_orig_indices = sort_perm[first_in_sorted] - unique_jdata = jdata[unique_orig_indices] - - num_lists = pixels_jt.num_tensors - if single_list: - unique_batch_idx = torch.zeros(num_unique, dtype=torch.long, device=device) - else: - unique_batch_idx = jidx.to(torch.long)[unique_orig_indices] - counts_per_list = torch.bincount(unique_batch_idx, minlength=num_lists) - new_offsets = torch.zeros(num_lists + 1, dtype=torch.long, device=device) - new_offsets[1:] = counts_per_list.cumsum(0) - - unique_pixels = JaggedTensor.from_data_and_offsets(unique_jdata, new_offsets) - return unique_pixels, inverse_indices, True - - def _sparse_render_impl( + def _render_sparse( self, pixels_to_render: JaggedTensor, - w2c: torch.Tensor, - K: torch.Tensor, - W: int, - H: int, + world_to_camera_matrices: torch.Tensor, + projection_matrices: torch.Tensor, + image_width: int, + image_height: int, near: float, far: float, camera_model: CameraModel, @@ -1940,72 +1689,43 @@ def _sparse_render_impl( sh_degree_to_use: int, tile_size: int, min_radius_2d: float, - eps2d: float, + eps_2d: float, antialias: bool, backgrounds: torch.Tensor | None, masks: torch.Tensor | None, - include_colors: bool, - include_depth: bool, - ) -> tuple[torch.Tensor, torch.Tensor]: - """Common implementation for all sparse_render_* methods. - - Returns ``(rendered_features_jdata, rendered_alphas_jdata)`` in the - *original* (possibly duplicated) pixel ordering. - """ - unique_pixels, inverse_indices, has_duplicates = self._deduplicate_pixels(pixels_to_render, W, H) - render_pixels = unique_pixels if has_duplicates else pixels_to_render - - C = w2c.size(0) - radii, means2d, depths, conics, compensations = self._do_projection( - w2c, - K, - W, - H, - eps2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, + render_mode: GaussianRenderMode, + ) -> tuple[JaggedTensor, JaggedTensor]: + """All four stages for an arbitrary set of pixels, in the requested pixel order.""" + projected, opacities = self._project_and_opacities( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, ) - 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, conics, opacities, C, tile_size, W, H) - - rendered_jdata, alphas_jdata = self._rasterize_screen_space_sparse( - render_pixels, - means2d, - conics, - features, - opacities, - W, - H, - tile_size, - tile_offsets, - tile_gaussian_ids, - active_tiles, - tile_pixel_mask, - tile_pixel_cumsum, - pixel_map, - backgrounds, - masks, + features = self._features(projected, world_to_camera_matrices, sh_degree_to_use, render_mode) + sparse_tiles = intersect_gaussian_tiles_sparse( + pixels_to_render, projected, tile_size=tile_size, opacities=opacities + ) + return rasterize_screen_space_gaussians_sparse( + projected, features, opacities, sparse_tiles, backgrounds=backgrounds, tile_masks=masks ) - if has_duplicates: - rendered_jdata = rendered_jdata.index_select(0, inverse_indices) - alphas_jdata = alphas_jdata.index_select(0, inverse_indices) - - return rendered_jdata, alphas_jdata + @staticmethod + def _sparse_result( + pixels_to_render: JaggedTensor | torch.Tensor, features: JaggedTensor, alphas: JaggedTensor + ) -> tuple[Any, Any]: + """Return sparse results as JaggedTensors, or as ``[C, P, D]`` tensors for tensor pixel input.""" + if isinstance(pixels_to_render, torch.Tensor): + return torch.stack(features.unbind(), dim=0), torch.stack(alphas.unbind(), dim=0) + return features, alphas def project_gaussians_for_depths( self, @@ -2024,7 +1744,7 @@ def project_gaussians_for_depths( ) -> ProjectedGaussianSplats: """ Projects this :class:`GaussianSplat3d` onto one or more image planes for rendering depth images in those planes. - You can render depth images from the projected Gaussians by calling :meth:`render_projected_gaussians`. + You can render depth images from the projected Gaussians by calling :meth:`render_from_projected_gaussians`. .. note:: @@ -2066,7 +1786,7 @@ def project_gaussians_for_depths( crop_origin_h=10) # To get the depth images, divide the last channel by the alpha values - true_depths_1 = cropped_images_1[..., -1:] / cropped_alphas + true_depths_1 = cropped_depth_images_1[..., -1:] / cropped_alphas Args: world_to_camera_matrices (torch.Tensor): Tensor of shape ``(C, 4, 4)`` representing the world-to-camera transformation matrices for ``C`` cameras. @@ -2096,42 +1816,21 @@ def project_gaussians_for_depths( This object contains the projected 2D representations of the Gaussians, which can be used for rendering depth images or further processing. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = depths.unsqueeze(-1) - opacities = self._make_opacities(C, compensations, antialias) - return ProjectedGaussianSplats( - radii=radii, - means2d=means2d, - depths=depths, - conics=conics, - compensations=compensations, - render_quantities=render_features, - opacities=opacities, + return self._project_for( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, image_width=image_width, image_height=image_height, - antialias=antialias, - eps_2d=eps_2d, - near_plane=near, - far_plane=far, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, sh_degree_to_use=-1, - camera_model=camera_model, - projection_method=self._resolve_projection_method(camera_model, projection_method), - _private=ProjectedGaussianSplats.__PRIVATE__, + render_mode=GaussianRenderMode.DEPTH, ) def project_gaussians_for_images( @@ -2152,7 +1851,7 @@ def project_gaussians_for_images( ) -> ProjectedGaussianSplats: """ Projects this :class:`GaussianSplat3d` onto one or more image planes for rendering multi-channel (see :attr:`num_channels`) images in those planes. - You can render images from the projected Gaussians by calling :meth:`render_projected_gaussians`. + You can render images from the projected Gaussians by calling :meth:`render_from_projected_gaussians`. .. note:: @@ -2224,42 +1923,21 @@ def project_gaussians_for_images( This object contains the projected 2D representations of the Gaussians, which can be used for rendering images or further processing. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = self._eval_sh(world_to_camera_matrices, radii, sh_degree_to_use) - opacities = self._make_opacities(C, compensations, antialias) - return ProjectedGaussianSplats( - radii=radii, - means2d=means2d, - depths=depths, - conics=conics, - compensations=compensations, - render_quantities=render_features, - opacities=opacities, + return self._project_for( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, image_width=image_width, image_height=image_height, - antialias=antialias, - eps_2d=eps_2d, - near_plane=near, - far_plane=far, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, sh_degree_to_use=sh_degree_to_use, - camera_model=camera_model, - projection_method=self._resolve_projection_method(camera_model, projection_method), - _private=ProjectedGaussianSplats.__PRIVATE__, + render_mode=GaussianRenderMode.FEATURES, ) def project_gaussians_for_images_and_depths( @@ -2281,7 +1959,7 @@ def project_gaussians_for_images_and_depths( """ Projects this :class:`GaussianSplat3d` onto one or more image planes for rendering multi-channel (see :attr:`num_channels`) images with depths in the last channel. - You can render images+depths from the projected Gaussians by calling :meth:`render_projected_gaussians`. + You can render images+depths from the projected Gaussians by calling :meth:`render_from_projected_gaussians`. .. note:: @@ -2314,13 +1992,16 @@ def project_gaussians_for_images_and_depths( # in each image plane. # Returns a tensor of shape [C, 100, 100, D] containing the images (where D is num_channels + 1 for depth), # and a tensor of shape [C, 100, 100, 1] containing the final alpha (opacity) values - # of each pixel. + # of each pixel. Binning the Gaussians into tiles once and passing the result lets several + # crops share it. + tiles = projected_gaussians.tile_intersection() cropped_images_1, cropped_alphas = gaussian_splat_3d.render_from_projected_gaussians( projected_gaussians, crop_width=100, crop_height=100, crop_origin_w=10, - crop_origin_h=10) + crop_origin_h=10, + tiles=tiles) cropped_images = cropped_images_1[..., :-1] # Extract image channels @@ -2358,49 +2039,21 @@ def project_gaussians_for_images_and_depths( This object contains the projected 2D representations of the Gaussians, which can be used for rendering images or further processing. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = self._make_render_features( - world_to_camera_matrices, - radii, - depths, - sh_degree_to_use, - include_colors=True, - include_depth=True, - ) - opacities = self._make_opacities(C, compensations, antialias) - return ProjectedGaussianSplats( - radii=radii, - means2d=means2d, - depths=depths, - conics=conics, - compensations=compensations, - render_quantities=render_features, - opacities=opacities, + return self._project_for( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, image_width=image_width, image_height=image_height, - antialias=antialias, - eps_2d=eps_2d, - near_plane=near, - far_plane=far, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, sh_degree_to_use=sh_degree_to_use, - camera_model=camera_model, - projection_method=self._resolve_projection_method(camera_model, projection_method), - _private=ProjectedGaussianSplats.__PRIVATE__, + render_mode=GaussianRenderMode.FEATURES_AND_DEPTH, ) def render_from_projected_gaussians( @@ -2410,9 +2063,10 @@ def render_from_projected_gaussians( crop_height: int = -1, crop_origin_w: int = -1, crop_origin_h: int = -1, - tile_size: int = 16, + tile_size: int | None = None, backgrounds: torch.Tensor | None = None, masks: torch.Tensor | None = None, + tiles: GaussianTileIntersection | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ Render a set of images from Gaussian splats that have already been projected onto image planes @@ -2422,14 +2076,17 @@ def render_from_projected_gaussians( .. note:: - If you want to render the full image, pass negative values for ``crop_width``, ``crop_height``, - ``crop_origin_w``, and ``crop_origin_h`` (default behavior). To render full images, - all these values must be negative or this method will raise an error. + A negative value for any of ``crop_width``, ``crop_height``, ``crop_origin_w`` and ``crop_origin_h`` + means its default: the full image width or height, or an origin of zero. All four negative (the + default) renders the full image. .. note:: - If your crop goes beyond the image boundaries, the resulting image will be clipped to - be within the image boundaries. + The output always has the requested crop size. Where the crop runs past the image boundary, + the part inside the image is rendered and the rest is filled with the background at zero + alpha; a crop that lies entirely outside the image is all background. The stage function + returns the clipped part only (empty for a crop entirely outside); this method pads it. Either + way the output is differentiable with respect to the projection whenever the projection is. Example: @@ -2450,13 +2107,16 @@ def render_from_projected_gaussians( # in each image plane. # Returns a tensor of shape [C, 100, 100, D] containing the images (where D is num_channels + 1 for depth), # and a tensor of shape [C, 100, 100, 1] containing the final alpha (opacity) values - # of each pixel. + # of each pixel. Binning the Gaussians into tiles once and passing the result lets several + # crops share it. + tiles = projected_gaussians.tile_intersection() cropped_images_1, cropped_alphas = gaussian_splat_3d.render_from_projected_gaussians( projected_gaussians, crop_width=100, crop_height=100, crop_origin_w=10, - crop_origin_h=10) + crop_origin_h=10, + tiles=tiles) cropped_images = cropped_images_1[..., :-1] # Extract image channels @@ -2470,21 +2130,27 @@ def render_from_projected_gaussians( :meth:`project_gaussians_for_images`, :meth:`project_gaussians_for_depths`, :meth:`project_gaussians_for_images_and_depths`, etc. crop_width (int): The width of the crop to render. If -1, the full image width is used. - Default is -1. + Default is -1. A crop that runs past the image edge is filled with the background at zero + alpha outside the image, so the output always has the requested size. crop_height (int): The height of the crop to render. If -1, the full image height is used. Default is -1. crop_origin_w (int): The x-coordinate of the top-left corner of the crop. If -1, the crop starts at (0, 0). Default is -1. crop_origin_h (int): The y-coordinate of the top-left corner of the crop. If -1, the crop starts at (0, 0). Default is -1. - tile_size (int): The size of the tiles to use for rendering. Default is 16. - This parameter controls the size of the tiles used for rendering the images. - You shouldn't set this parameter unless you really know what you are doing. + tile_size (int | None): The size of the tiles to use for rendering. Taken from ``tiles`` when + they are given, and 16 otherwise; passing both raises if they disagree. You shouldn't set + this parameter unless you really know what you are doing. backgrounds (torch.Tensor | None): Optional background colors of shape ``(C, D)``. If ``None``, background is treated as 0. - masks (torch.Tensor | None): Optional per-pixel boolean mask of shape ``(C, cropH, cropW)`` - (in crop coordinate space, matching the output dimensions). - ``True`` means render, ``False`` means skip (filled with background). + masks (torch.Tensor | None): Optional per-pixel boolean mask in crop coordinates, of the requested + crop size ``(C, cropH, cropW)`` or of its size after clipping to the image, on the projection's + device. ``True`` means render, ``False`` means skip (filled with background). Without a crop it + is a full-image mask. + tiles (GaussianTileIntersection | None): The tile intersections of ``projected_gaussians`` at + ``tile_size``, from :meth:`ProjectedGaussianSplats.tile_intersection`. Computed here when + ``None``. Pass them when rendering several crops from one projection, so the Gaussians are + binned into tiles once rather than once per crop. Returns: @@ -2497,81 +2163,35 @@ def render_from_projected_gaussians( and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ pg = projected_gaussians - W = pg.image_width - H = pg.image_height - C = pg.radii.size(0) - - raster_w = crop_width if crop_width > 0 else W - raster_h = crop_height if crop_height > 0 else H + projected = pg.projected_gaussians + width, height = projected.image_width, projected.image_height + crop_w = crop_width if crop_width >= 0 else width + crop_h = crop_height if crop_height >= 0 else height origin_w = crop_origin_w if crop_origin_w >= 0 else 0 origin_h = crop_origin_h if crop_origin_h >= 0 else 0 - is_crop = raster_w != W or raster_h != H or origin_w != 0 or origin_h != 0 - - tile_masks = _pixel_mask_to_tile_mask(masks, tile_size) if masks is not None else None - - if is_crop: - num_tiles_h = math.ceil(raster_h / tile_size) - num_tiles_w = math.ceil(raster_w / tile_size) - tile_offsets, tile_gaussian_ids = _C.intersect_gaussian_tiles( - pg.means2d, - pg.radii, - pg.depths, - C, - tile_size, - num_tiles_h, - num_tiles_w, - conics=pg.inv_covar_2d, - opacities=pg.opacities, - ) - features, alphas = cast( - tuple[torch.Tensor, torch.Tensor], - _RasterizeScreenSpaceGaussiansFn.apply( - pg.means2d, - pg.inv_covar_2d, - pg.render_quantities, - pg.opacities, - raster_w, - raster_h, - origin_w, - origin_h, - tile_size, - tile_offsets, - tile_gaussian_ids, - False, - backgrounds, - tile_masks, - ), - ) - else: - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - pg.means2d, - pg.radii, - pg.depths, - pg.inv_covar_2d, - pg.opacities, - C, - tile_size, - W, - H, - ) - features, alphas = self._rasterize_screen_space( - pg.means2d, - pg.inv_covar_2d, - pg.render_quantities, - pg.opacities, - W, - H, - tile_size, - tile_offsets, - tile_gaussian_ids, - backgrounds, - tile_masks, - ) - - if masks is not None: - features, alphas = _apply_pixel_mask(features, alphas, masks, backgrounds) - - return features, alphas + is_crop = crop_w != width or crop_h != height or origin_w != 0 or origin_h != 0 + requested_h, requested_w = crop_h, crop_w + if tiles is not None: + if tile_size is not None and tiles.tile_size != tile_size: + raise ValueError(f"tiles were computed at tile_size {tiles.tile_size}, not the requested {tile_size}") + tile_size = tiles.tile_size + elif tile_size is None: + tile_size = 16 + # The stage function clips the crop at the image edge (to nothing, if it lies entirely outside), + # checks the crop mask against the requested or clipped crop size, and returns the clipped part; it + # is padded back below, so the output has the requested size. + crop = (origin_w, origin_h, crop_w, crop_h) if is_crop else None + images, alphas = rasterize_screen_space_gaussians( + projected, + pg.render_quantities, + pg.opacities, + tiles if tiles is not None else pg.tile_intersection(tile_size), + backgrounds=backgrounds, + masks=masks if crop is None else None, + crop=crop, + crop_masks=masks if crop is not None else None, + ) + return pad_crop(images, alphas, requested_h, requested_w, backgrounds) def render_depths( self, @@ -2655,51 +2275,26 @@ def render_depths( Each element represents the alpha value (opacity) at a pixel such that ``0 <= alpha < 1``, and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = depths.unsqueeze(-1) - opacities = self._make_opacities(C, compensations, antialias) - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - means2d, - radii, - depths, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - tile_masks = _pixel_mask_to_tile_mask(masks, tile_size) if masks is not None else None - features, alphas = self._rasterize_screen_space( - means2d, - conics, - render_features, - opacities, - image_width, - image_height, - tile_size, - tile_offsets, - tile_gaussian_ids, - backgrounds, - tile_masks, + return self._render_dense( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=-1, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.DEPTH, + world_space=False, ) - if masks is not None: - features, alphas = _apply_pixel_mask(features, alphas, masks, backgrounds) - return features, alphas def sparse_render_depths( self, @@ -2785,41 +2380,28 @@ def sparse_render_depths( and ``P`` is the number of pixel coordinates rendered per camera. Each element represents the alpha value (opacity) at that pixel such that ``0 <= alpha < 1``, and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ - if isinstance(pixels_to_render, torch.Tensor): - pixels_jt = JaggedTensor(impl=JaggedTensorCpp(pixels_to_render)) - elif isinstance(pixels_to_render, JaggedTensor): - pixels_jt = pixels_to_render - else: - raise TypeError("pixels_to_render must be either a torch.Tensor or a fvdb.JaggedTensor") - - rendered_jdata, alphas_jdata = self._sparse_render_impl( - pixels_jt, - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - near, - far, - camera_model, - projection_method, - distortion_coeffs, - -1, - tile_size, - min_radius_2d, - eps_2d, - antialias, - backgrounds, - masks, - include_colors=False, - include_depth=True, + pixels_jt = as_pixel_jagged(pixels_to_render) + features, alphas = self._render_sparse( + pixels_to_render=pixels_jt, + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=-1, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.DEPTH, ) - ret_features = pixels_jt.jagged_like(rendered_jdata) - ret_alphas = pixels_jt.jagged_like(alphas_jdata) - - if isinstance(pixels_to_render, torch.Tensor): - return ret_features._impl.jdata, ret_alphas._impl.jdata - else: - return ret_features, ret_alphas + return self._sparse_result(pixels_to_render, features, alphas) def render_images( self, @@ -2905,51 +2487,26 @@ def render_images( Each element represents the alpha value (opacity) at a pixel such that ``0 <= alpha < 1``, and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = self._eval_sh(world_to_camera_matrices, radii, sh_degree_to_use) - opacities = self._make_opacities(C, compensations, antialias) - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - means2d, - radii, - depths, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - tile_masks = _pixel_mask_to_tile_mask(masks, tile_size) if masks is not None else None - features, alphas = self._rasterize_screen_space( - means2d, - conics, - render_features, - opacities, - image_width, - image_height, - tile_size, - tile_offsets, - tile_gaussian_ids, - backgrounds, - tile_masks, + return self._render_dense( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=sh_degree_to_use, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.FEATURES, + world_space=False, ) - if masks is not None: - features, alphas = _apply_pixel_mask(features, alphas, masks, backgrounds) - return features, alphas def render_images_from_world( self, @@ -2969,6 +2526,8 @@ def render_images_from_world( antialias: bool = False, backgrounds: torch.Tensor | None = None, masks: torch.Tensor | None = None, + crop: Crop | None = None, + crop_masks: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ Render dense images by rasterizing directly from world-space 3D Gaussians. @@ -3040,62 +2599,38 @@ def render_images_from_world( masks (torch.Tensor | None): Optional per-pixel boolean mask of shape ``(C, H, W)``. ``True`` means render, ``False`` means skip (filled with background). + crop (tuple[int, int, int, int] | None): Optional ``(origin_w, origin_h, width, height)`` window + to render instead of the full image, clipped to the image; tiles outside it are skipped and + the output has the clipped size; a crop entirely outside the image gives an empty render. + crop_masks (torch.Tensor | None): Optional per-pixel boolean mask in crop coordinates, of the crop's + requested or clipped size, as an alternative to the image-coordinate ``masks`` when a crop is given. Returns: images (torch.Tensor): Rendered images of shape ``(C, H, W, D)``. alpha_images (torch.Tensor): Alpha images of shape ``(C, H, W, 1)``. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = self._eval_sh(world_to_camera_matrices, radii, sh_degree_to_use) - opacities = self._make_opacities(C, compensations, antialias) - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - means2d, - radii, - depths, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - tile_masks = _pixel_mask_to_tile_mask(masks, tile_size) if masks is not None else None - if distortion_coeffs is None: - distortion_coeffs = torch.zeros( - C, 12, device=world_to_camera_matrices.device, dtype=world_to_camera_matrices.dtype - ) - features, alphas = self._rasterize_world_space( - render_features, - opacities, - world_to_camera_matrices, - projection_matrices, - distortion_coeffs, - camera_model, - image_width, - image_height, - tile_size, - tile_offsets, - tile_gaussian_ids, - backgrounds, - tile_masks, + return self._render_dense( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=sh_degree_to_use, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.FEATURES, + world_space=True, + crop=crop, + crop_masks=crop_masks, ) - if masks is not None: - features, alphas = _apply_pixel_mask(features, alphas, masks, backgrounds) - return features, alphas def render_depths_from_world( self, @@ -3114,64 +2649,37 @@ def render_depths_from_world( antialias: bool = False, backgrounds: torch.Tensor | None = None, masks: torch.Tensor | None = None, + crop: Crop | None = None, + crop_masks: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ Render dense depth images by rasterizing directly from world-space 3D Gaussians. This mirrors :meth:`render_images_from_world`, but renders depth-only outputs with the - same camera-model and projection-method dispatch. + same camera-model and projection-method dispatch. ``crop`` and ``crop_masks`` behave as there. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = depths.unsqueeze(-1) - opacities = self._make_opacities(C, compensations, antialias) - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - means2d, - radii, - depths, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - tile_masks = _pixel_mask_to_tile_mask(masks, tile_size) if masks is not None else None - if distortion_coeffs is None: - distortion_coeffs = torch.zeros( - C, 12, device=world_to_camera_matrices.device, dtype=world_to_camera_matrices.dtype - ) - features, alphas = self._rasterize_world_space( - render_features, - opacities, - world_to_camera_matrices, - projection_matrices, - distortion_coeffs, - camera_model, - image_width, - image_height, - tile_size, - tile_offsets, - tile_gaussian_ids, - backgrounds, - tile_masks, + return self._render_dense( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=-1, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.DEPTH, + world_space=True, + crop=crop, + crop_masks=crop_masks, ) - if masks is not None: - features, alphas = _apply_pixel_mask(features, alphas, masks, backgrounds) - return features, alphas def sparse_render_images( self, @@ -3261,41 +2769,28 @@ def sparse_render_images( Each element represents the alpha value (opacity) at that pixel such that ``0 <= alpha < 1``, and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ - if isinstance(pixels_to_render, torch.Tensor): - pixels_jt = JaggedTensor(impl=JaggedTensorCpp(pixels_to_render)) - elif isinstance(pixels_to_render, JaggedTensor): - pixels_jt = pixels_to_render - else: - raise TypeError("pixels_to_render must be either a torch.Tensor or a fvdb.JaggedTensor") - - rendered_jdata, alphas_jdata = self._sparse_render_impl( - pixels_jt, - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - near, - far, - camera_model, - projection_method, - distortion_coeffs, - sh_degree_to_use, - tile_size, - min_radius_2d, - eps_2d, - antialias, - backgrounds, - masks, - include_colors=True, - include_depth=False, + pixels_jt = as_pixel_jagged(pixels_to_render) + features, alphas = self._render_sparse( + pixels_to_render=pixels_jt, + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=sh_degree_to_use, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.FEATURES, ) - ret_features = pixels_jt.jagged_like(rendered_jdata) - ret_alphas = pixels_jt.jagged_like(alphas_jdata) - - if isinstance(pixels_to_render, torch.Tensor): - return ret_features._impl.jdata, ret_alphas._impl.jdata - else: - return ret_features, ret_alphas + return self._sparse_result(pixels_to_render, features, alphas) def sparse_render_images_and_depths( self, @@ -3387,41 +2882,28 @@ def sparse_render_images_and_depths( Each element represents the alpha value (opacity) at that pixel such that ``0 <= alpha < 1``, and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ - if isinstance(pixels_to_render, torch.Tensor): - pixels_jt = JaggedTensor(impl=JaggedTensorCpp(pixels_to_render)) - elif isinstance(pixels_to_render, JaggedTensor): - pixels_jt = pixels_to_render - else: - raise TypeError("pixels_to_render must be either a torch.Tensor or a fvdb.JaggedTensor") - - rendered_jdata, alphas_jdata = self._sparse_render_impl( - pixels_jt, - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - near, - far, - camera_model, - projection_method, - distortion_coeffs, - sh_degree_to_use, - tile_size, - min_radius_2d, - eps_2d, - antialias, - backgrounds, - masks, - include_colors=True, - include_depth=True, + pixels_jt = as_pixel_jagged(pixels_to_render) + features, alphas = self._render_sparse( + pixels_to_render=pixels_jt, + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=sh_degree_to_use, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.FEATURES_AND_DEPTH, ) - ret_features = pixels_jt.jagged_like(rendered_jdata) - ret_alphas = pixels_jt.jagged_like(alphas_jdata) - - if isinstance(pixels_to_render, torch.Tensor): - return ret_features._impl.jdata, ret_alphas._impl.jdata - else: - return ret_features, ret_alphas + return self._sparse_result(pixels_to_render, features, alphas) def render_images_and_depths( self, @@ -3510,58 +2992,26 @@ def render_images_and_depths( Each element represents the alpha value (opacity) at a pixel such that ``0 <= alpha < 1``, and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = self._make_render_features( - world_to_camera_matrices, - radii, - depths, - sh_degree_to_use, - include_colors=True, - include_depth=True, - ) - opacities = self._make_opacities(C, compensations, antialias) - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - means2d, - radii, - depths, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - tile_masks = _pixel_mask_to_tile_mask(masks, tile_size) if masks is not None else None - features, alphas = self._rasterize_screen_space( - means2d, - conics, - render_features, - opacities, - image_width, - image_height, - tile_size, - tile_offsets, - tile_gaussian_ids, - backgrounds, - tile_masks, + return self._render_dense( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=sh_degree_to_use, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.FEATURES_AND_DEPTH, + world_space=False, ) - if masks is not None: - features, alphas = _apply_pixel_mask(features, alphas, masks, backgrounds) - return features, alphas def render_images_and_depths_from_world( self, @@ -3581,71 +3031,38 @@ def render_images_and_depths_from_world( antialias: bool = False, backgrounds: torch.Tensor | None = None, masks: torch.Tensor | None = None, + crop: Crop | None = None, + crop_masks: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: """ Render dense RGBD images by rasterizing directly from world-space 3D Gaussians. This mirrors :meth:`render_images_from_world`, but returns image channels with depth in the - final channel while using the same camera-model and projection-method dispatch. + final channel while using the same camera-model and projection-method dispatch. ``crop`` and + ``crop_masks`` behave as there. """ - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - render_features = self._make_render_features( - world_to_camera_matrices, - radii, - depths, - sh_degree_to_use, - include_colors=True, - include_depth=True, - ) - opacities = self._make_opacities(C, compensations, antialias) - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - means2d, - radii, - depths, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - tile_masks = _pixel_mask_to_tile_mask(masks, tile_size) if masks is not None else None - if distortion_coeffs is None: - distortion_coeffs = torch.zeros( - C, 12, device=world_to_camera_matrices.device, dtype=world_to_camera_matrices.dtype - ) - features, alphas = self._rasterize_world_space( - render_features, - opacities, - world_to_camera_matrices, - projection_matrices, - distortion_coeffs, - camera_model, - image_width, - image_height, - tile_size, - tile_offsets, - tile_gaussian_ids, - backgrounds, - tile_masks, + return self._render_dense( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + sh_degree_to_use=sh_degree_to_use, + tile_size=tile_size, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, + backgrounds=backgrounds, + masks=masks, + render_mode=GaussianRenderMode.FEATURES_AND_DEPTH, + world_space=True, + crop=crop, + crop_masks=crop_masks, ) - if masks is not None: - features, alphas = _apply_pixel_mask(features, alphas, masks, backgrounds) - return features, alphas def render_num_contributing_gaussians( self, @@ -3687,7 +3104,7 @@ def render_num_contributing_gaussians( near, # near clipping plane far) # far clipping plane - num_gaussians_cij = num_gaussians[c, i, j, 0] # Number of contributing Gaussians at pixel (i, j) in camera c + num_gaussians_cij = num_gaussians[c, i, j] # Number of contributing Gaussians at pixel (i, j) in camera c Args: world_to_camera_matrices (torch.Tensor): Tensor of shape ``(C, 4, 4)`` representing the @@ -3715,54 +3132,31 @@ def render_num_contributing_gaussians( antialias (bool): If ``True``, applies opacity correction to the projected Gaussians when using ``eps_2d > 0.0``. Returns: - images (torch.Tensor): A tensor of shape ``(C, H, W, 1)`` where ``C`` is the number of camera views, - ``H`` is the height of the images, ``W`` is the width of the images. - Each element represents the number of contributing Gaussians at that pixel. - alpha_images (torch.Tensor): A tensor of shape ``(C, H, W, 1)`` where ``C`` is the number of camera views, + num_contributing (torch.Tensor): An ``int32`` tensor of shape ``(C, H, W)`` where ``C`` is the number of + camera views, ``H`` is the height of the images, ``W`` is the width of the images. + Each element is the number of contributing Gaussians at that pixel. + alpha_images (torch.Tensor): A tensor of shape ``(C, H, W)`` where ``C`` is the number of camera views, ``H`` is the height of the images, and ``W`` is the width of the images. Each element represents the alpha value (opacity) at a pixel such that ``0 <= alpha < 1``, and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ with torch.no_grad(): - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - opacities = self._make_opacities(C, compensations, antialias) - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - means2d, - radii, - depths, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - return _C.rasterize_num_contributing_gaussians( - means2d, - conics, - opacities, - tile_offsets, - tile_gaussian_ids, - image_width, - image_height, - 0, - 0, - tile_size, + projected, opacities = self._project_and_opacities( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, ) + tiles = intersect_gaussian_tiles(projected, tile_size=tile_size, opacities=opacities) + return rasterize_num_contributing_gaussians(projected, opacities, tiles) @overload def sparse_render_num_contributing_gaussians( @@ -3830,7 +3224,7 @@ def sparse_render_num_contributing_gaussians( Args: pixels_to_render (torch.Tensor | JaggedTensor): A :class:`fvdb.JaggedTensor` of shape ``(C, R_c, 2)`` representing the pixels to render for each camera, where ``C`` is the number of camera views and ``R_c`` is the - number of pixels to render per camera. Each value is an (x, y) pixel coordinate. + number of pixels to render per camera. Each value is a ``(row, col)`` pixel coordinate. world_to_camera_matrices (torch.Tensor): Tensor of shape ``(C, 4, 4)`` representing the world-to-camera transformation matrices for C cameras. Each matrix transforms points from world coordinates to camera coordinates. @@ -3866,82 +3260,27 @@ def sparse_render_num_contributing_gaussians( Each element represents the alpha value (opacity) at that pixel such that ``0 <= alpha < 1``, and 0 means the pixel is fully transparent, and 1 means the pixel is fully opaque. """ - is_dense = isinstance(pixels_to_render, torch.Tensor) - if is_dense: - C, R, _ = pixels_to_render.shape - tensors = [pixels_to_render[i] for i in range(C)] - pixels_jt = JaggedTensor(tensors) - else: - pixels_jt = pixels_to_render - + pixels_jt = as_pixel_jagged(pixels_to_render) with torch.no_grad(): - unique_pixels_jt, inverse_indices, has_dups = self._deduplicate_pixels(pixels_jt, image_width, image_height) - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - 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, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - result_ncg, result_alphas = _C.rasterize_num_contributing_gaussians_sparse( - means2d, - conics, - opacities, - tile_offsets, - tile_gaussian_ids, - unique_pixels_jt._impl, - active_tiles, - tile_pixel_mask, - tile_pixel_cumsum, - pixel_map, - image_width, - image_height, - 0, - 0, - tile_size, + projected, opacities = self._project_and_opacities( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, ) - - ncg_jt = JaggedTensor(impl=result_ncg) - alphas_jt = JaggedTensor(impl=result_alphas) - - if has_dups: - ncg_jt = pixels_jt.jagged_like(ncg_jt.jdata.index_select(0, inverse_indices)) - alphas_jt = pixels_jt.jagged_like(alphas_jt.jdata.index_select(0, inverse_indices)) - - if is_dense: - return ( - torch.stack(ncg_jt.unbind(), dim=0), - torch.stack(alphas_jt.unbind(), dim=0), + sparse_tiles = intersect_gaussian_tiles_sparse( + pixels_jt, projected, tile_size=tile_size, opacities=opacities ) - return ncg_jt, alphas_jt + counts, alphas = rasterize_num_contributing_gaussians_sparse(projected, opacities, sparse_tiles) + return self._sparse_result(pixels_to_render, counts, alphas) def render_contributing_gaussian_ids( self, @@ -3997,70 +3336,23 @@ def render_contributing_gaussian_ids( jagged tensor containing the weights of the contributing Gaussians of each rendered pixel for each camera. The weights are in row-major order and sum to 1 for each pixel if that pixel is opaque (alpha=1). """ - # TODO: Projection currently always evaluates SH, but this method only needs - # geometric projection (2D means, conics, opacities) -- the SH color values are - # unused. Ideally rendering should be more generic: accept an arbitrary feature - # tensor (e.g. integer IDs, raw features) without requiring SH evaluation. That - # would also let us avoid the wasted SH computation here and support additional - # shading models in the future. For now we just render "deep IDs" as a fixed - # function. (Ported from the C++ renderContributingGaussianIdsImpl TODO.) with torch.no_grad(): - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, - ) - C = world_to_camera_matrices.size(0) - opacities = self._make_opacities(C, compensations, antialias) - tile_offsets, tile_gaussian_ids, _, _ = self._intersect_tiles( - means2d, - radii, - depths, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - ncg = None - if top_k_contributors <= 0: - ncg, _ = _C.rasterize_num_contributing_gaussians( - means2d, - conics, - opacities, - tile_offsets, - tile_gaussian_ids, - image_width, - image_height, - 0, - 0, - tile_size, - ) - ids, weights = _C.rasterize_contributing_gaussian_ids( - means2d, - conics, - opacities, - tile_offsets, - tile_gaussian_ids, - image_width, - image_height, - 0, - 0, - tile_size, - top_k_contributors, - ncg, + projected, opacities = self._project_and_opacities( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, ) - return JaggedTensor(impl=ids), JaggedTensor(impl=weights) + tiles = intersect_gaussian_tiles(projected, tile_size=tile_size, opacities=opacities) + return rasterize_contributing_gaussian_ids(projected, opacities, tiles, top_k_contributors) @overload def sparse_render_contributing_gaussian_ids( @@ -4130,7 +3422,7 @@ def sparse_render_contributing_gaussian_ids( pixels_to_render (torch.Tensor | JaggedTensor): A :class:`torch.Tensor` of shape ``(C, R, 2)`` or a :class:`fvdb.JaggedTensor` of shape ``(C, R_c, 2)`` representing the pixels to render for each camera, where ``C`` is the number of camera views and ``R``/``R_c`` is the - number of pixels to render per camera. Each value is an (x, y) pixel coordinate. + number of pixels to render per camera. Each value is a ``(row, col)`` pixel coordinate. world_to_camera_matrices (torch.Tensor): Tensor of shape ``(C, 4, 4)`` representing the world-to-camera transformation matrices for ``C`` cameras. Each matrix transforms points from world coordinates to camera coordinates. @@ -4163,146 +3455,26 @@ def sparse_render_contributing_gaussian_ids( weights (fvdb.JaggedTensor): A ``[[C1P1 + C1P2 + ... C1PN1, 1], ... [CNP1 + CNP2 + ... CNPNN, 1]]`` jagged tensor containing the weights of the contributing Gaussians of each rendered pixel for each camera. The weights are in row-major order and sum to 1 for each pixel if that pixel is opaque (alpha=1). """ - if isinstance(pixels_to_render, torch.Tensor): - C, R, _ = pixels_to_render.shape - tensors = [pixels_to_render[i] for i in range(C)] - pixels_jt = JaggedTensor(tensors) - else: - pixels_jt = pixels_to_render - + pixels_jt = as_pixel_jagged(pixels_to_render) with torch.no_grad(): - unique_pixels_jt, inverse_indices, has_dups = self._deduplicate_pixels(pixels_jt, image_width, image_height) - radii, means2d, depths, conics, compensations = self._do_projection( - world_to_camera_matrices, - projection_matrices, - image_width, - image_height, - eps_2d, - near, - far, - min_radius_2d, - antialias, - camera_model, - projection_method, - distortion_coeffs, + projected, opacities = self._project_and_opacities( + world_to_camera_matrices=world_to_camera_matrices, + projection_matrices=projection_matrices, + image_width=image_width, + image_height=image_height, + near=near, + far=far, + camera_model=camera_model, + projection_method=projection_method, + distortion_coeffs=distortion_coeffs, + min_radius_2d=min_radius_2d, + eps_2d=eps_2d, + antialias=antialias, ) - 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, - conics, - opacities, - C, - tile_size, - image_width, - image_height, - ) - ncg_jt = None - if top_k_contributors <= 0: - ncg_jt, _ = _C.rasterize_num_contributing_gaussians_sparse( - means2d, - conics, - opacities, - tile_offsets, - tile_gaussian_ids, - unique_pixels_jt._impl, - active_tiles, - tile_pixel_mask, - tile_pixel_cumsum, - pixel_map, - image_width, - image_height, - 0, - 0, - tile_size, - ) - ids, weights = _C.rasterize_contributing_gaussian_ids_sparse( - means2d, - conics, - opacities, - tile_offsets, - tile_gaussian_ids, - unique_pixels_jt._impl, - active_tiles, - tile_pixel_mask, - tile_pixel_cumsum, - pixel_map, - image_width, - image_height, - 0, - 0, - tile_size, - top_k_contributors, - ncg_jt, + sparse_tiles = intersect_gaussian_tiles_sparse( + pixels_jt, projected, tile_size=tile_size, opacities=opacities ) - ids_jt = JaggedTensor(impl=ids) - weights_jt = JaggedTensor(impl=weights) - if has_dups: - # `ids`/`weights` are CONTRIBUTION-major: one row per (pixel, contributor), - # with each pixel owning a variable-length segment. They therefore cannot be - # indexed by pixel the way a per-pixel result can (see sparse_render, which - # does exactly that on a one-row-per-pixel array and is correct). A duplicated - # pixel needs a copy of its unique pixel's whole segment. - # - # Both tensors share the same jagged structure, so the index arithmetic is - # computed once and applied to each payload. - plan = self._contribution_expansion_plan(ids_jt, pixels_jt, inverse_indices) - ids_jt = self._apply_contribution_expansion(plan, ids_jt) - weights_jt = self._apply_contribution_expansion(plan, weights_jt) - return ids_jt, weights_jt - - @staticmethod - def _contribution_expansion_plan( - unique_jt: JaggedTensor, - pixels_jt: JaggedTensor, - inverse_indices: torch.Tensor, - ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Index arithmetic for re-expanding a per-(pixel, contributor) result. - - Maps the deduplicated pixel list back onto the caller's pixel list by repeating - each unique pixel's whole contribution segment. Depends only on the jagged - structure, not on the payload, so callers with several equally-shaped tensors - (ids and weights) build one plan and apply it to each. - - Returns ``(gather, offsets, list_ids)`` for - :meth:`JaggedTensor.from_data_offsets_and_list_ids`. - """ - device = unique_jt.jdata.device - offsets_unique = unique_jt.joffsets.to(device) - starts = offsets_unique[inverse_indices] # segment start per output pixel - counts = offsets_unique[1:][inverse_indices] - starts # segment length per output pixel - - offsets = torch.zeros(counts.numel() + 1, dtype=torch.long, device=device) - offsets[1:] = counts.cumsum(0) - - segment = torch.repeat_interleave(torch.arange(counts.numel(), device=device), counts) - within = torch.arange(int(offsets[-1]), device=device) - offsets[segment] - gather = starts[segment] + within - - # Rebuild (camera, pixel-within-camera) ids from the ORIGINAL pixel list. - camera = pixels_jt.jidx.to(device).long() - within_camera = torch.arange(camera.numel(), device=device) - pixels_jt.joffsets.to(device)[camera] - list_ids = torch.stack([camera, within_camera], dim=1).to(torch.int32) - return gather, offsets, list_ids - - @staticmethod - def _apply_contribution_expansion( - plan: tuple[torch.Tensor, torch.Tensor, torch.Tensor], - jt: JaggedTensor, - ) -> JaggedTensor: - """Apply a plan from :meth:`_contribution_expansion_plan` to one payload.""" - gather, offsets, list_ids = plan - return JaggedTensor.from_data_offsets_and_list_ids(jt.jdata.index_select(0, gather), offsets, list_ids) + return rasterize_contributing_gaussian_ids_sparse(projected, opacities, sparse_tiles, top_k_contributors) def relocate_gaussians( self, @@ -4326,13 +3498,8 @@ def relocate_gaussians( Returns: tuple[torch.Tensor, torch.Tensor]: Tuple of (logit_opacities_new [N], log_scales_new [N, 3]). """ - return _C.mcmc_relocate_gaussians( - log_scales, - logit_opacities, - ratios, - binomial_coeffs, - n_max, - min_opacity, + return fvdb_functional.mcmc_relocate_gaussians( + log_scales, logit_opacities, ratios, binomial_coeffs, n_max, min_opacity ) def add_noise_to_means(self, noise_scale: float, t: float = 0.005, k: float = 100.0) -> None: @@ -4344,14 +3511,8 @@ def add_noise_to_means(self, noise_scale: float, t: float = 0.005, k: float = 10 t (float): Parameter t for noise scaling. Defaults to 0.005. k (float): Parameter k for noise scaling. Defaults to 100.0. """ - _C.mcmc_add_noise_to_means( - self._means, - self._log_scales, - self._logit_opacities, - self._quats, - noise_scale, - t, - k, + fvdb_functional.mcmc_add_noise_to_means( + self._means, self._log_scales, self._logit_opacities, self._quats, noise_scale, t, k ) def reset_accumulated_gradient_state(self) -> None: @@ -4391,15 +3552,8 @@ def save_ply( """ if isinstance(filename, pathlib.Path): filename = str(filename) - _C.save_gaussian_ply( - filename, - self._means, - self._quats, - self._log_scales, - self._logit_opacities, - self._sh0, - self._shN, - metadata, + fvdb_functional.save_gaussian_ply( + filename, self._means, self._quats, self._log_scales, self._logit_opacities, self._sh0, self._shN, metadata ) @overload @@ -4638,32 +3792,6 @@ def state_dict(self) -> dict[str, torch.Tensor]: d["accumulated_max_2d_radii"] = self._accumulated_max_2d_radii return d - @staticmethod - def _camera_model_from_cpp(camera_model: _C.CameraModel) -> CameraModel: - try: - return CameraModel[camera_model.name] - except KeyError as exc: - raise ValueError(f"Invalid camera model: {camera_model}") from exc - - @staticmethod - def _camera_model_to_cpp(camera_model: CameraModel) -> _C.CameraModel: - if isinstance(camera_model, CameraModel): - return getattr(_C.CameraModel, camera_model.name) - return camera_model - - @staticmethod - def _projection_method_from_cpp(projection_method: _C.ProjectionMethod) -> ProjectionMethod: - try: - return ProjectionMethod[projection_method.name] - except KeyError as exc: - raise ValueError(f"Invalid projection method: {projection_method}") from exc - - @staticmethod - def _projection_method_to_cpp(projection_method: ProjectionMethod) -> _C.ProjectionMethod: - if isinstance(projection_method, ProjectionMethod): - return getattr(_C.ProjectionMethod, projection_method.name) - return projection_method - # TODO: Make a batched class to encapsulate this jagged rendering pipeline. def gaussian_render_jagged( @@ -4826,7 +3954,7 @@ def gaussian_render_jagged( # --- Non-differentiable tile intersection --- 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( + tile_offsets, tile_gaussian_ids_t = fvdb_functional.intersect_gaussian_tiles( means2d, radii, depths, diff --git a/pyproject.toml b/pyproject.toml index b0fccbe3..077a8852 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -22,7 +22,7 @@ authors = [ license="Apache-2.0" license-files = ["LICEN[CS]E*"] dependencies = [ - "fvdb-core>=0.6.0dev0", + "fvdb-core>=0.7.0dev0", "boto3", "dlnr_lite", "nanovdb-editor<0.2.0", diff --git a/tests/benchmarks/test_3dgs.py b/tests/benchmarks/test_3dgs.py index 9f1119b6..5d0e579f 100644 --- a/tests/benchmarks/test_3dgs.py +++ b/tests/benchmarks/test_3dgs.py @@ -14,6 +14,7 @@ import torch.utils.data import yaml +import fvdb_reality_capture.functional as frc_functional from fvdb_reality_capture import CameraModel # Set multiprocessing start method to 'spawn' to avoid fork() warnings with PyTorch @@ -160,15 +161,14 @@ def run_project_gaussians(self): ) def run_render_gaussians(self): - # Render an image from the gaussian splats - # possibly using a crop of the full image - self.colors, self.alphas = self.runner.model.render_from_projected_gaussians( - self.projected_gaussians, - crop_width=self.image_width, - crop_height=self.image_height, - crop_origin_w=0, - crop_origin_h=0, - tile_size=self.runner.config.tile_size, + # Render the full image from the projected Gaussians through the stage functions, so tile + # intersection is timed as part of rendering on every iteration. + pg = self.projected_gaussians + tiles = frc_functional.intersect_gaussian_tiles( + pg.projected_gaussians, pg.opacities, tile_size=self.runner.config.tile_size + ) + self.colors, self.alphas = frc_functional.rasterize_screen_space_gaussians( + pg.projected_gaussians, pg.render_quantities, pg.opacities, tiles ) def run_forward(self): diff --git a/tests/unit/test_functional_gaussian_splatting.py b/tests/unit/test_functional_gaussian_splatting.py new file mode 100644 index 00000000..066ef825 --- /dev/null +++ b/tests/unit/test_functional_gaussian_splatting.py @@ -0,0 +1,674 @@ +# Copyright Contributors to the OpenVDB Project +# SPDX-License-Identifier: Apache-2.0 +# +"""Tests for the composable Gaussian splatting pipeline in ``fvdb_reality_capture.functional``. + +The pipeline is checked against :class:`GaussianSplat3d`, which composes the same stages, and +against itself across the dense, sparse, cropped and world-space paths. +""" + +import unittest + +import numpy as np +import torch +from fvdb import JaggedTensor +from fvdb.utils.tests import get_fvdb_test_data_path + +import fvdb_reality_capture.functional as F +from fvdb_reality_capture import CameraModel, GaussianRenderMode, GaussianSplat3d, ProjectionMethod + + +def rgb_to_sh(rgb: torch.Tensor) -> torch.Tensor: + C0 = 0.28209479177387814 + return (rgb - 0.5) / C0 + + +class FunctionalPipelineTestCase(unittest.TestCase): + """Load the garden scene once per test and expose its tensors.""" + + def setUp(self): + if not torch.cuda.is_available(): + self.skipTest("Gaussian splatting requires a CUDA device") + torch.manual_seed(0) + np.random.seed(0) + self.device = torch.device("cuda:0") + data = np.load(get_fvdb_test_data_path() / "gsplat" / "test_garden_cropped.npz") + + self.means = torch.from_numpy(data["means3d"]).float().to(self.device) + self.quats = torch.from_numpy(data["quats"]).float().to(self.device) + self.log_scales = torch.log(torch.from_numpy(data["scales"]).float().to(self.device)) + self.logit_opacities = torch.logit(torch.from_numpy(data["opacities"]).float().to(self.device)) + colors = torch.from_numpy(data["colors"]).float().to(self.device) + self.W = int(data["width"].item()) + self.H = int(data["height"].item()) + all_w2c = torch.from_numpy(data["viewmats"]).float().to(self.device) + all_K = torch.from_numpy(data["Ks"]).float().to(self.device) + self.w2c = all_w2c[:2].contiguous() + self.K = all_K[:2].contiguous() + self.C = self.w2c.shape[0] + + N = self.means.shape[0] + self.sh_degree = 2 + sh = torch.zeros(N, (self.sh_degree + 1) ** 2, 3, device=self.device) + sh[:, 0] = rgb_to_sh(colors) + sh[:, 1:] = torch.randn_like(sh[:, 1:]) * 0.05 + self.sh0 = sh[:, :1].contiguous() + self.shN = sh[:, 1:].contiguous() + + def _params(self, requires_grad: bool = False): + tensors = [self.means, self.quats, self.log_scales, self.logit_opacities, self.sh0, self.shN] + return [t.detach().clone().requires_grad_(requires_grad) for t in tensors] + + def _model(self, params) -> GaussianSplat3d: + means, quats, log_scales, logit_opacities, sh0, shN = params + return GaussianSplat3d.from_tensors( + means=means, quats=quats, log_scales=log_scales, logit_opacities=logit_opacities, sh0=sh0, shN=shN + ) + + def _render_functional(self, params, render_mode=GaussianRenderMode.FEATURES, **kwargs): + means, quats, log_scales, logit_opacities, sh0, shN = params + projected = F.project_gaussians(means, quats, log_scales, self.w2c, self.K, self.W, self.H, **kwargs) + features = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected, render_mode=render_mode) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + tiles = F.intersect_gaussian_tiles(projected, opacities) + return F.rasterize_screen_space_gaussians(projected, features, opacities, tiles) + + +class TestStageOutputs(FunctionalPipelineTestCase): + def test_projection_contract(self): + params = self._params() + projected = F.project_gaussians(*params[:3], self.w2c, self.K, self.W, self.H, antialias=True) + C, N = self.C, self.means.shape[0] + self.assertEqual(tuple(projected.radii.shape), (C, N, 2)) + self.assertEqual(tuple(projected.means2d.shape), (C, N, 2)) + self.assertEqual(tuple(projected.depths.shape), (C, N)) + self.assertEqual(tuple(projected.conics.shape), (C, N, 3)) + self.assertEqual(tuple(projected.compensations.shape), (C, N)) + self.assertEqual((projected.num_cameras, projected.num_gaussians), (C, N)) + self.assertEqual(projected.projection_method, ProjectionMethod.ANALYTIC) + self.assertTrue(projected.is_differentiable) + with self.assertRaises(Exception): + projected.radii = None # frozen + # Equality is identity, so the dataclasses can be compared and hashed despite holding tensors. + self.assertEqual(projected, projected) + self.assertNotEqual(projected, without := F.project_gaussians(*params[:3], self.w2c, self.K, self.W, self.H)) + self.assertEqual(len({projected, without}), 2) + + without = F.project_gaussians(*params[:3], self.w2c, self.K, self.W, self.H, antialias=False) + self.assertIsNone(without.compensations) + + def test_render_modes(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + projected = F.project_gaussians(means, quats, log_scales, self.w2c, self.K, self.W, self.H) + features = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected) + depth = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected, render_mode=GaussianRenderMode.DEPTH) + both = F.evaluate_gaussian_sh( + means, sh0, shN, self.w2c, projected, render_mode=GaussianRenderMode.FEATURES_AND_DEPTH + ) + self.assertEqual(features.shape[-1], 3) + self.assertEqual(depth.shape[-1], 1) + self.assertEqual(both.shape[-1], 4) + torch.testing.assert_close(both[..., :3], features) + torch.testing.assert_close(both[..., 3:], depth) + # Both the projection and the features stage zero the depth of culled Gaussians and agree on the rest. + visible = (projected.radii > 0).all(-1) + torch.testing.assert_close(depth[..., 0][visible], projected.depths[visible]) + self.assertTrue(bool((projected.depths[~visible] == 0).all())) + # Lower degrees are allowed, higher than available are not. + F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected, sh_degree_to_use=0) + with self.assertRaises(ValueError): + F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected, sh_degree_to_use=self.sh_degree + 1) + + def test_tile_intersection_contract(self): + params = self._params() + projected = F.project_gaussians(*params[:3], self.w2c, self.K, self.W, self.H) + tiles = F.intersect_gaussian_tiles(projected, F.compute_gaussian_opacities(params[3], projected), tile_size=16) + self.assertEqual(tuple(tiles.tile_offsets.shape), (self.C, tiles.num_tiles_h, tiles.num_tiles_w)) + self.assertEqual(tiles.num_tiles_h, -(-self.H // 16)) + self.assertEqual(tiles.num_tiles_w, -(-self.W // 16)) + # Culling with opacities can only remove intersections relative to the bounding-box test. + loose = F.intersect_gaussian_tiles(projected, None, tile_size=16) + self.assertLessEqual(tiles.tile_gaussian_ids.numel(), loose.tile_gaussian_ids.numel()) + + +class TestMatchesGaussianSplat3d(FunctionalPipelineTestCase): + def test_forward_matches_oo(self): + params = self._params() + images, alphas = self._render_functional(params) + images_oo, alphas_oo = self._model(params).render_images(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + torch.testing.assert_close(images, images_oo, atol=1e-5, rtol=1e-5) + torch.testing.assert_close(alphas, alphas_oo, atol=1e-5, rtol=1e-5) + + def test_backward_matches_oo(self): + params_fn = self._params(requires_grad=True) + params_oo = self._params(requires_grad=True) + images, alphas = self._render_functional(params_fn) + (images.square().mean() + alphas.mean()).backward() + images_oo, alphas_oo = self._model(params_oo).render_images(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + (images_oo.square().mean() + alphas_oo.mean()).backward() + for p_fn, p_oo in zip(params_fn, params_oo): + self.assertIsNotNone(p_fn.grad) + self.assertGreater(float(p_fn.grad.abs().max()), 0.0) + # Both paths run the same atomicAdd kernels, so agreement is up to summation order. + torch.testing.assert_close(p_fn.grad, p_oo.grad, atol=1e-5, rtol=1e-4) + + def test_depth_and_features_and_depth_match_oo(self): + params = self._params() + model = self._model(params) + depth, alpha = self._render_functional(params, render_mode=GaussianRenderMode.DEPTH) + depth_oo, alpha_oo = model.render_depths(self.w2c, self.K, self.W, self.H, 0.01, 1e10, min_radius_2d=0.0) + torch.testing.assert_close(depth, depth_oo, atol=1e-5, rtol=1e-5) + both, _ = self._render_functional(params, render_mode=GaussianRenderMode.FEATURES_AND_DEPTH) + both_oo, _ = model.render_images_and_depths(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + torch.testing.assert_close(both, both_oo, atol=1e-5, rtol=1e-5) + + def test_world_space_matches_oo_and_is_differentiable_through_ut(self): + params = self._params(requires_grad=True) + means, quats, log_scales, logit_opacities, sh0, shN = params + projected = F.project_gaussians( + means, quats, log_scales, self.w2c, self.K, self.W, self.H, projection_method=ProjectionMethod.UNSCENTED + ) + self.assertFalse(projected.is_differentiable) + self.assertFalse(projected.means2d.requires_grad) + features = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + tiles = F.intersect_gaussian_tiles(projected, opacities) + images, alphas = F.rasterize_world_space_gaussians( + means, quats, log_scales, projected, features, opacities, self.w2c, self.K, tiles + ) + images.mean().backward() + for p in (means, quats, log_scales, logit_opacities, sh0): + self.assertGreater(float(p.grad.abs().max()), 0.0) + + # OpenCV models need explicit distortion coefficients; zeros would silently render pinhole rays. + opencv = F.project_gaussians( + means.detach(), + quats.detach(), + log_scales.detach(), + self.w2c, + self.K, + self.W, + self.H, + camera_model=CameraModel.OPENCV_RADTAN_5, + distortion_coeffs=torch.zeros(self.C, 12, device=self.device), + ) + with self.assertRaises(RuntimeError): + F.rasterize_world_space_gaussians( + means, quats, log_scales, opencv, features.detach(), opacities.detach(), self.w2c, self.K, tiles + ) + # The coefficients are checked before the kernel reads twelve per camera. + world_args = (means, quats, log_scales, opencv, features.detach(), opacities.detach(), self.w2c, self.K, tiles) + with self.assertRaisesRegex(RuntimeError, "shape"): + F.rasterize_world_space_gaussians(*world_args, distortion_coeffs=torch.zeros(self.C, 5, device=self.device)) + with self.assertRaisesRegex(RuntimeError, "must be on"): + F.rasterize_world_space_gaussians(*world_args, distortion_coeffs=torch.zeros(self.C, 12)) + # Pinhole cameras ignore the coefficients, whatever is passed. + F.rasterize_world_space_gaussians( + means, + quats, + log_scales, + projected, + features.detach(), + opacities.detach(), + self.w2c, + self.K, + tiles, + distortion_coeffs=torch.zeros(self.C, 5), + ) + + params_oo = self._params() + images_oo, alphas_oo = self._model(params_oo).render_images_from_world( + self.w2c, self.K, self.W, self.H, 0.01, 1e10, projection_method=ProjectionMethod.UNSCENTED + ) + torch.testing.assert_close(images.detach(), images_oo, atol=1e-5, rtol=1e-5) + torch.testing.assert_close(alphas.detach(), alphas_oo, atol=1e-5, rtol=1e-5) + + def test_depth_is_differentiable_through_the_unscented_projection(self): + means, quats, log_scales, logit_opacities, sh0, shN = self._params(requires_grad=True) + analytic = F.project_gaussians( + means.detach(), quats.detach(), log_scales.detach(), self.w2c, self.K, self.W, self.H + ) + projected = F.project_gaussians( + means, quats, log_scales, self.w2c, self.K, self.W, self.H, projection_method=ProjectionMethod.UNSCENTED + ) + self.assertFalse(projected.depths.requires_grad) + # Depth is linear in the center, so the recomputed depth matches the projection's exactly and + # carries the gradient the unscented kernel cannot provide. + depth = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected, render_mode=GaussianRenderMode.DEPTH) + self.assertTrue(depth.requires_grad) + # Each projection zeroes the depth of what it culls, and the two cull slightly differently. + visible = (analytic.radii.amin(-1) > 0) & (projected.radii.amin(-1) > 0) + torch.testing.assert_close(depth[..., 0][visible], analytic.depths[visible], atol=1e-4, rtol=1e-4) + depth.sum().backward() + self.assertGreater(float(means.grad.abs().max()), 0.0) + both = F.evaluate_gaussian_sh( + means, sh0, shN, self.w2c, projected, render_mode=GaussianRenderMode.FEATURES_AND_DEPTH + ) + self.assertTrue(both[..., 3:].requires_grad) + # The depth channel is the same function of its inputs with or without autograd. + with torch.no_grad(): + same = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected, render_mode=GaussianRenderMode.DEPTH) + self.assertTrue(torch.equal(same, both[..., 3:].detach())) + visible = (projected.radii > 0).all(-1) + torch.testing.assert_close(same[..., 0][visible], projected.depths[visible]) + self.assertFalse(F.requires_distortion_coeffs(CameraModel.PINHOLE)) + self.assertTrue(F.requires_distortion_coeffs(CameraModel.OPENCV_RADTAN_5)) + + def test_training_loop_reduces_loss(self): + target, _ = self._render_functional(self._params()) + params = self._params(requires_grad=True) + means, quats, log_scales, logit_opacities, sh0, shN = params + with torch.no_grad(): + sh0.add_(0.3 * torch.randn_like(sh0)) + optimizer = torch.optim.Adam([sh0, shN, logit_opacities], lr=1e-2) + losses = [] + for _ in range(15): + optimizer.zero_grad() + images, _ = self._render_functional(params) + loss = torch.nn.functional.l1_loss(images, target) + loss.backward() + optimizer.step() + losses.append(loss.item()) + self.assertLess(losses[-1], 0.5 * losses[0]) + + +class TestSparseAndCrop(FunctionalPipelineTestCase): + def _pixels(self, with_duplicates: bool) -> JaggedTensor: + rows = torch.randint(0, self.H, (200,), device=self.device) + cols = torch.randint(0, self.W, (200,), device=self.device) + px = torch.stack([rows, cols], dim=1) + px = torch.unique(px, dim=0) + if with_duplicates: + px = torch.cat([px, px[:37], px[5:6].repeat(3, 1)]) + return JaggedTensor([px, px.flip(0)]) + + def test_sparse_matches_dense_at_requested_pixels(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + projected = F.project_gaussians(means, quats, log_scales, self.w2c, self.K, self.W, self.H) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + features = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected) + dense, dense_alpha = F.rasterize_screen_space_gaussians( + projected, features, opacities, F.intersect_gaussian_tiles(projected, opacities) + ) + for with_duplicates in (False, True): + pixels = self._pixels(with_duplicates) + sparse_tiles = F.intersect_gaussian_tiles_sparse(pixels, projected, opacities) + self.assertEqual(sparse_tiles.has_duplicates, with_duplicates) + rendered, alphas = F.rasterize_screen_space_gaussians_sparse(projected, features, opacities, sparse_tiles) + self.assertEqual(len(rendered), self.C) + all_tiles = torch.ones_like(sparse_tiles.active_tile_mask) + same, _ = F.rasterize_screen_space_gaussians_sparse( + projected, features, opacities, sparse_tiles, tile_masks=all_tiles + ) + torch.testing.assert_close(same.jdata, rendered.jdata) + # A per-pixel mask handed in as a tile mask is refused instead of skipping the wrong tiles. + with self.assertRaisesRegex(ValueError, "per-tile"): + F.rasterize_screen_space_gaussians_sparse( + projected, features, opacities, sparse_tiles, tile_masks=torch.ones(self.C, self.H, self.W) + ) + for c in range(self.C): + px = pixels[c].jdata + self.assertEqual(tuple(rendered[c].jdata.shape), (px.shape[0], 3)) + torch.testing.assert_close(rendered[c].jdata, dense[c, px[:, 0], px[:, 1]], atol=1e-5, rtol=1e-5) + torch.testing.assert_close(alphas[c].jdata, dense_alpha[c, px[:, 0], px[:, 1]], atol=1e-5, rtol=1e-5) + + def test_sparse_matches_oo_and_backward_runs(self): + params = self._params(requires_grad=True) + means, quats, log_scales, logit_opacities, sh0, shN = params + pixels = self._pixels(with_duplicates=True) + projected = F.project_gaussians(means, quats, log_scales, self.w2c, self.K, self.W, self.H) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + features = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected) + sparse_tiles = F.intersect_gaussian_tiles_sparse(pixels, projected, opacities) + rendered, alphas = F.rasterize_screen_space_gaussians_sparse(projected, features, opacities, sparse_tiles) + rendered_oo, alphas_oo = self._model(self._params()).sparse_render_images( + pixels, self.w2c, self.K, self.W, self.H, 0.01, 1e10 + ) + torch.testing.assert_close(rendered.jdata.detach(), rendered_oo.jdata, atol=1e-5, rtol=1e-5) + torch.testing.assert_close(alphas.jdata.detach(), alphas_oo.jdata, atol=1e-5, rtol=1e-5) + rendered.jdata.mean().backward() + self.assertGreater(float(means.grad.abs().max()), 0.0) + self.assertGreater(float(sh0.grad.abs().max()), 0.0) + + def test_crop_is_a_slice_of_the_full_render(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + projected = F.project_gaussians(means, quats, log_scales, self.w2c, self.K, self.W, self.H) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + features = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected) + tiles = F.intersect_gaussian_tiles(projected, opacities) + full, full_alpha = F.rasterize_screen_space_gaussians(projected, features, opacities, tiles) + ox, oy, w, h = 37, 21, 90, 60 + crop, crop_alpha = F.rasterize_screen_space_gaussians( + projected, features, opacities, tiles, crop=(ox, oy, w, h) + ) + self.assertEqual(tuple(crop.shape), (self.C, h, w, 3)) + torch.testing.assert_close(crop, full[:, oy : oy + h, ox : ox + w], atol=1e-5, rtol=1e-5) + torch.testing.assert_close(crop_alpha, full_alpha[:, oy : oy + h, ox : ox + w], atol=1e-5, rtol=1e-5) + # Clamped to the image, invalid crops rejected. + clamped, _ = F.rasterize_screen_space_gaussians( + projected, features, opacities, tiles, crop=(self.W - 10, self.H - 5, 100, 100) + ) + self.assertEqual(tuple(clamped.shape[1:3]), (5, 10)) + for bad in ((-1, 0, 10, 10), (0, 0, 0, 10)): + with self.assertRaises(ValueError): + F.rasterize_screen_space_gaussians(projected, features, opacities, tiles, crop=bad) + # A crop entirely outside the image clips to nothing rather than raising. + empty, _ = F.rasterize_screen_space_gaussians(projected, features, opacities, tiles, crop=(self.W, 0, 10, 10)) + self.assertEqual(tuple(empty.shape), (self.C, 0, 0, 3)) + + # The OO crop path agrees, including a mask given in crop coordinates. + model = self._model(params) + pg = model.project_gaussians_for_images(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + crop_oo, _ = model.render_from_projected_gaussians( + pg, crop_width=w, crop_height=h, crop_origin_w=ox, crop_origin_h=oy + ) + torch.testing.assert_close(crop_oo, crop, atol=1e-5, rtol=1e-5) + mask = torch.zeros(self.C, h, w, dtype=torch.bool, device=self.device) + mask[:, : h // 2] = True + masked, masked_alpha = model.render_from_projected_gaussians( + pg, crop_width=w, crop_height=h, crop_origin_w=ox, crop_origin_h=oy, masks=mask + ) + torch.testing.assert_close(masked[:, : h // 2], crop[:, : h // 2], atol=1e-5, rtol=1e-5) + self.assertEqual(float(masked[:, h // 2 :].abs().max()), 0.0) + self.assertEqual(float(masked_alpha[:, h // 2 :].abs().max()), 0.0) + # A crop past the image edge keeps its requested size: the part inside the image is the render, + # the rest is background with zero alpha. A crop-space mask of the requested size still applies. + edge_w, edge_h = 50, 40 + edge_mask = torch.ones(self.C, edge_h, edge_w, dtype=torch.bool, device=self.device) + edge, edge_alpha = model.render_from_projected_gaussians( + pg, + crop_width=edge_w, + crop_height=edge_h, + crop_origin_w=self.W - 30, + crop_origin_h=self.H - 25, + masks=edge_mask, + ) + self.assertEqual(tuple(edge.shape[1:3]), (edge_h, edge_w)) + self.assertEqual(tuple(edge_alpha.shape[1:]), (edge_h, edge_w, 1)) + torch.testing.assert_close(edge[:, :25, :30], full[:, self.H - 25 :, self.W - 30 :], atol=1e-5, rtol=1e-5) + self.assertEqual(float(edge[:, 25:].abs().max()), 0.0) + self.assertEqual(float(edge[:, :, 30:].abs().max()), 0.0) + self.assertEqual(float(edge_alpha[:, 25:].abs().max()), 0.0) + self.assertEqual(float(edge_alpha[:, :, 30:].abs().max()), 0.0) + # Non-boolean masks are accepted with and without a crop. + float_mask = torch.ones(self.C, self.H, self.W, device=self.device) + with_float, _ = F.rasterize_screen_space_gaussians( + projected, features, opacities, tiles, masks=float_mask, crop=(ox, oy, w, h) + ) + torch.testing.assert_close(with_float, crop, atol=1e-5, rtol=1e-5) + # With a crop the stage also takes a crop-space mask through crop_masks, and pools the skipped + # tiles from it; masks stays in image coordinates, so neither shape is ever guessed. + stage_masked, _ = F.rasterize_screen_space_gaussians( + projected, features, opacities, tiles, crop=(ox, oy, w, h), crop_masks=mask + ) + torch.testing.assert_close(stage_masked, masked, atol=1e-5, rtol=1e-5) + with self.assertRaisesRegex(ValueError, "not both"): + F.rasterize_screen_space_gaussians( + projected, features, opacities, tiles, masks=float_mask, crop=(ox, oy, w, h), crop_masks=mask + ) + with self.assertRaisesRegex(ValueError, "needs a crop"): + F.rasterize_screen_space_gaussians(projected, features, opacities, tiles, crop_masks=mask) + # A crop mask of any other size is rejected rather than read from its top-left corner; the class + # method hands its crop-space mask to the same check. + with self.assertRaisesRegex(ValueError, "crop_masks must match the crop .* clipped size"): + F.rasterize_screen_space_gaussians( + projected, features, opacities, tiles, crop=(ox, oy, w, h), crop_masks=mask[:, :-1] + ) + with self.assertRaisesRegex(ValueError, "clipped size"): + model.render_from_projected_gaussians( + pg, crop_width=w, crop_height=h, crop_origin_w=ox, crop_origin_h=oy, masks=mask[:, :-1] + ) + # A mask on the wrong device is rejected at the same point, not deep inside torch. + cpu_full_mask = torch.ones(self.C, self.H, self.W, dtype=torch.bool) + with self.assertRaisesRegex(ValueError, "must be on"): + F.rasterize_screen_space_gaussians(projected, features, opacities, tiles, masks=cpu_full_mask) + with self.assertRaisesRegex(ValueError, "must be on"): + model.render_from_projected_gaussians( + pg, crop_width=w, crop_height=h, crop_origin_w=ox, crop_origin_h=oy, masks=mask.cpu() + ) + + def test_analysis_matches_oo(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + model = self._model(params) + projected = F.project_gaussians(means, quats, log_scales, self.w2c, self.K, self.W, self.H) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + tiles = F.intersect_gaussian_tiles(projected, opacities) + counts, alphas = F.rasterize_num_contributing_gaussians(projected, opacities, tiles) + counts_oo, alphas_oo = model.render_num_contributing_gaussians(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + self.assertTrue(torch.equal(counts, counts_oo)) + torch.testing.assert_close(alphas, alphas_oo) + + ids, weights = F.rasterize_contributing_gaussian_ids(projected, opacities, tiles) + ids_oo, weights_oo = model.render_contributing_gaussian_ids(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + self.assertTrue(torch.equal(ids.jdata, ids_oo.jdata)) + self.assertEqual(ids.ldim, 2) + self.assertEqual(int(ids.jdata.numel()), int(counts.sum())) + top_ids, _ = F.rasterize_contributing_gaussian_ids(projected, opacities, tiles, top_k_contributors=3) + self.assertLessEqual(int(top_ids.jdata.numel()), int(counts.clamp(max=3).sum())) + + def test_sparse_analysis_with_duplicates(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + projected = F.project_gaussians(means, quats, log_scales, self.w2c, self.K, self.W, self.H) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + tiles = F.intersect_gaussian_tiles(projected, opacities) + counts_dense, _ = F.rasterize_num_contributing_gaussians(projected, opacities, tiles) + ids_dense, _ = F.rasterize_contributing_gaussian_ids(projected, opacities, tiles) + + pixels = self._pixels(with_duplicates=True) + sparse_tiles = F.intersect_gaussian_tiles_sparse(pixels, projected, opacities) + counts, _ = F.rasterize_num_contributing_gaussians_sparse(projected, opacities, sparse_tiles) + ids, weights = F.rasterize_contributing_gaussian_ids_sparse(projected, opacities, sparse_tiles) + self.assertEqual(ids.ldim, 2) + for c in range(self.C): + px = pixels[c].jdata + self.assertTrue(torch.equal(counts[c].jdata, counts_dense[c, px[:, 0], px[:, 1]])) + + def flat(x): + return x.jdata if isinstance(x, JaggedTensor) else x + + expected = [flat(ids_dense[c][int(r) * self.W + int(q)]) for r, q in px.tolist()] + got = [flat(t) for t in ids[c].unbind()] + self.assertEqual(len(got), len(expected)) + for g, e in zip(got, expected): + self.assertTrue(torch.equal(g, e)) + + def test_sparse_analysis_keeps_a_trailing_camera_with_no_requested_pixels(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + w2c, K = self.w2c[:2], self.K[:2] + projected = F.project_gaussians(means, quats, log_scales, w2c, K, self.W, self.H) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + busy = torch.tensor([[self.H // 2, self.W // 2], [self.H // 2 + 3, self.W // 2 + 3]], device=self.device) + empty = torch.empty(0, 2, dtype=torch.int64, device=self.device) + # Camera 0 has a duplicate (so the expansion path runs); camera 1 requests nothing. + pixels = JaggedTensor([torch.cat([busy, busy[:1]]), empty]) + sparse_tiles = F.intersect_gaussian_tiles_sparse(pixels, projected, opacities) + self.assertTrue(sparse_tiles.has_duplicates) + for result in F.rasterize_contributing_gaussian_ids_sparse(projected, opacities, sparse_tiles): + self.assertEqual(len(result), 2) + self.assertEqual(len(result[0].unbind()), 3) + self.assertEqual(len(result[1].unbind()), 0) + counts, _ = F.rasterize_num_contributing_gaussians_sparse(projected, opacities, sparse_tiles) + self.assertEqual(len(counts), 2) + self.assertEqual(counts[1].jdata.numel(), 0) + + def test_single_camera_sparse_analysis_keeps_camera_nesting(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + w2c, K = self.w2c[:1], self.K[:1] + projected = F.project_gaussians(means, quats, log_scales, w2c, K, self.W, self.H) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + pixels = self._pixels(with_duplicates=True)[0] # one camera, duplicates included + pixels = JaggedTensor([pixels.jdata]) + sparse_tiles = F.intersect_gaussian_tiles_sparse(pixels, projected, opacities) + ids, weights = F.rasterize_contributing_gaussian_ids_sparse(projected, opacities, sparse_tiles) + self.assertEqual(ids.ldim, 2) + self.assertEqual(len(ids), 1) + self.assertEqual(len(ids[0].unbind()), pixels.jdata.shape[0]) + counts, _ = F.rasterize_num_contributing_gaussians_sparse(projected, opacities, sparse_tiles) + self.assertEqual(int(ids.jdata.numel()), int(counts.jdata.sum())) + + def test_opacities_are_validated_and_made_contiguous(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + projected = F.project_gaussians( + means, quats, log_scales, self.w2c, self.K, image_width=self.W, image_height=self.H + ) + features = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + self.assertEqual(tuple(opacities.shape), (self.C, means.shape[0])) + tiles = F.intersect_gaussian_tiles(projected, opacities, tile_size=16) + a, _ = F.rasterize_screen_space_gaussians(projected, features, opacities, tiles) + # Every stage checks the opacities against the projection's camera and Gaussian counts. + with self.assertRaises(ValueError): + F.rasterize_screen_space_gaussians(projected, features, opacities[:1], tiles) + with self.assertRaises(ValueError): + F.intersect_gaussian_tiles(projected, opacities[:, :-1], tile_size=16) + with self.assertRaises(ValueError): + F.rasterize_num_contributing_gaussians(projected, torch.sigmoid(logit_opacities), tiles) + # A non-contiguous (expanded) opacity tensor is accepted and materialized, not handed to the kernel. + expanded = torch.sigmoid(logit_opacities).unsqueeze(0).expand(self.C, -1) + self.assertFalse(expanded.is_contiguous()) + c, _ = F.rasterize_screen_space_gaussians(projected, features, expanded, tiles) + torch.testing.assert_close(c, a) + + def test_projected_splats_snapshot_opacities_at_projection(self): + model = self._model(self._params(requires_grad=True)) + pg = model.project_gaussians_for_images(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + # Opacities and features come from the same projection, so they carry a gradient together. + self.assertTrue(pg.opacities.requires_grad) + self.assertTrue(pg.render_quantities.requires_grad) + expected = torch.sigmoid(model.logit_opacities.detach()).repeat(self.C, 1) + torch.testing.assert_close(pg.opacities.detach(), expected) + # An optimizer step between projection and render must not leak into the projection. + with torch.no_grad(): + model.logit_opacities.add_(1.0) + torch.testing.assert_close(pg.opacities.detach(), expected) + self.assertIs(pg.opacities, pg.opacities) + # A projection made without a graph renders without one, for both quantities. + with torch.no_grad(): + preview = model.project_gaussians_for_images(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + self.assertFalse(preview.opacities.requires_grad) + self.assertFalse(preview.render_quantities.requires_grad) + + def test_deduplicate_pixels_keeps_request_order_and_never_merges_out_of_image_pixels(self): + w, h = 64, 48 + # (0, 64) linearizes to the same key as (1, 0) in a 64-wide image; it must not be merged into it. + coords = torch.tensor([[3, 7], [1, 2], [3, 7], [0, 64], [1, 0], [5, 5], [1, 2]], dtype=torch.int64) + coords = coords.to(self.device) + unique, inverse, has_duplicates = F.deduplicate_pixels(JaggedTensor([coords]), w, h) + self.assertTrue(has_duplicates) + self.assertTrue(torch.equal(unique.jdata, coords[[0, 1, 3, 4, 5]]), "first occurrences, in request order") + self.assertTrue(torch.equal(unique.jdata[inverse], coords)) + + def test_as_pixel_jagged_validates_shape_and_dtype(self): + with self.assertRaises(TypeError): + F.as_pixel_jagged(torch.zeros(1, 4, 2, device=self.device)) + with self.assertRaises(ValueError): + F.as_pixel_jagged(torch.zeros(0, 4, 2, dtype=torch.int64, device=self.device)) + with self.assertRaises(ValueError): + F.as_pixel_jagged(JaggedTensor([torch.zeros(4, 3, dtype=torch.int64, device=self.device)])) + pixels = F.as_pixel_jagged(torch.zeros(self.C, 4, 2, dtype=torch.int32, device=self.device)) + self.assertEqual(pixels.num_tensors, self.C) + + def test_projected_splats_render_crops_outside_the_image_as_background(self): + model = self._model(self._params()) + pg = model.project_gaussians_for_images(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + self.assertIs(pg.opacities, pg.opacities) + # Tile intersections are computed on demand, not held by the projection. + self.assertIsNot(pg.tile_intersection(16), pg.tile_intersection(16)) + background = torch.tensor([[0.2, 0.4, 0.6]] * self.C, device=self.device) + for origin_w, origin_h in ((self.W, 0), (0, self.H), (self.W + 5, self.H + 5)): + images, alphas = model.render_from_projected_gaussians( + pg, + crop_width=10, + crop_height=10, + crop_origin_w=origin_w, + crop_origin_h=origin_h, + backgrounds=background, + ) + self.assertEqual(tuple(images.shape), (self.C, 10, 10, 3)) + torch.testing.assert_close(images, background[:, None, None, :].expand_as(images)) + self.assertEqual(float(alphas.abs().max()), 0.0) + + def test_crops_outside_the_image_stay_differentiable(self): + params = self._params(requires_grad=True) + model = self._model(params) + outside = dict(crop_width=10, crop_height=10, crop_origin_w=self.W, crop_origin_h=0) + pg = model.project_gaussians_for_images(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + images, alphas = model.render_from_projected_gaussians(pg, **outside) + self.assertTrue(images.requires_grad) + self.assertTrue(alphas.requires_grad) + # Nothing was rendered, so backward runs and every gradient it produces is zero. + (images.sum() + alphas.sum()).backward() + grads = [p.grad for p in params if p.grad is not None] + self.assertTrue(grads) + self.assertTrue(all(float(g.abs().max()) == 0.0 for g in grads)) + with torch.no_grad(): + frozen = model.project_gaussians_for_images(self.w2c, self.K, self.W, self.H, 0.01, 1e10) + images, alphas = model.render_from_projected_gaussians(frozen, **outside) + self.assertFalse(images.requires_grad) + self.assertFalse(alphas.requires_grad) + # The stage function has the same contract minus the padding: an all-outside crop is empty. + empty, empty_alpha = F.rasterize_screen_space_gaussians( + pg.projected_gaussians, + pg.render_quantities, + pg.opacities, + pg.tile_intersection(16), + crop=(self.W, 0, 10, 10), + ) + self.assertEqual(tuple(empty.shape), (self.C, 0, 0, 3)) + self.assertEqual(tuple(empty_alpha.shape), (self.C, 0, 0, 1)) + # ... and short-circuits the rasterizer while staying connected to the features. + self.assertTrue(empty.requires_grad) + # A mask of the wrong shape is rejected for an outside crop just as for any other crop. + bad_mask = torch.ones(self.C, 3, 3, dtype=torch.bool, device=self.device) + with self.assertRaisesRegex(ValueError, "clipped size"): + model.render_from_projected_gaussians(pg, masks=bad_mask, **outside) + good_mask = torch.ones(self.C, 10, 10, dtype=torch.bool, device=self.device) + images, _ = model.render_from_projected_gaussians(pg, masks=good_mask, **outside) + self.assertEqual(tuple(images.shape), (self.C, 10, 10, 3)) + # A crop mask the size of the image at a nonzero origin is read in crop coordinates, since crop_masks + # is a separate argument from the image-coordinate masks. + crop_mask = torch.zeros(self.C, self.H, self.W, dtype=torch.bool, device=self.device) + crop_mask[:, :20, :20] = True + shifted, _ = F.rasterize_screen_space_gaussians( + pg.projected_gaussians, + pg.render_quantities, + pg.opacities, + pg.tile_intersection(16), + crop=(5, 7, self.W, self.H), + crop_masks=crop_mask, + ) + plain, _ = F.rasterize_screen_space_gaussians( + pg.projected_gaussians, + pg.render_quantities, + pg.opacities, + pg.tile_intersection(16), + crop=(5, 7, self.W, self.H), + ) + torch.testing.assert_close(shifted[:, :20, :20], plain[:, :20, :20]) + self.assertEqual(float(shifted[:, 20:].detach().abs().max()), 0.0) + + def test_empty_selection(self): + params = self._params() + means, quats, log_scales, logit_opacities, sh0, shN = params + projected = F.project_gaussians(means, quats, log_scales, self.w2c, self.K, self.W, self.H) + opacities = F.compute_gaussian_opacities(logit_opacities, projected) + features = F.evaluate_gaussian_sh(means, sh0, shN, self.w2c, projected) + empty = JaggedTensor([torch.empty(0, 2, dtype=torch.int64, device=self.device) for _ in range(self.C)]) + sparse_tiles = F.intersect_gaussian_tiles_sparse(empty, projected, opacities) + rendered, alphas = F.rasterize_screen_space_gaussians_sparse(projected, features, opacities, sparse_tiles) + self.assertEqual(tuple(rendered.jdata.shape), (0, 3)) + self.assertEqual(tuple(alphas.jdata.shape), (0, 1)) + counts, _ = F.rasterize_num_contributing_gaussians_sparse(projected, opacities, sparse_tiles) + self.assertEqual(counts.jdata.numel(), 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/unit/test_gaussian_splat_3d.py b/tests/unit/test_gaussian_splat_3d.py index ff000878..dbee01bb 100644 --- a/tests/unit/test_gaussian_splat_3d.py +++ b/tests/unit/test_gaussian_splat_3d.py @@ -19,6 +19,7 @@ from parameterized import parameterized, parameterized_class from fvdb import Grid, JaggedTensor +import fvdb_reality_capture.functional as frc_functional from fvdb_reality_capture import ( CameraModel, GaussianSplat3d, @@ -5083,14 +5084,14 @@ def _render_jagged(cam_viewmats, cam_Ks): class TestDeduplicatePixels(unittest.TestCase): - """Unit tests for GaussianSplat3d._deduplicate_pixels (pure-Python implementation).""" + """Unit tests for fvdb_reality_capture.functional.deduplicate_pixels.""" IMAGE_WIDTH = 64 IMAGE_HEIGHT = 64 @staticmethod def _dedup(pixels_jt, w=64, h=64): - return GaussianSplat3d._deduplicate_pixels(pixels_jt, w, h) + return frc_functional.deduplicate_pixels(pixels_jt, w, h) @parameterized.expand([(torch.int32,), (torch.int64,)]) def test_empty(self, dtype): @@ -5115,7 +5116,7 @@ def test_all_unique(self, dtype): unique, inv, has_dups = self._dedup(pixels) self.assertFalse(has_dups) self.assertEqual(unique.jdata.shape[0], 5) - self.assertEqual(inv.shape[0], 5) + self.assertEqual(inv.shape[0], 0) # nothing to reorder without duplicates @parameterized.expand([(torch.int32,), (torch.int64,)]) def test_some_duplicates(self, dtype): diff --git a/tests/unit/test_gaussian_splat_api.py b/tests/unit/test_gaussian_splat_api.py index 0ca2b9f8..e7801bbd 100644 --- a/tests/unit/test_gaussian_splat_api.py +++ b/tests/unit/test_gaussian_splat_api.py @@ -20,27 +20,30 @@ def test_gaussian_splat_api_is_owned_by_reality_capture(): assert hasattr(fvdb_reality_capture, symbol) assert not hasattr(fvdb, symbol) + # The composable pipeline lives here; fvdb.functional only has the flat kernel wrappers. + for symbol in ("project_gaussians", "evaluate_gaussian_sh", "ProjectedGaussians"): + assert hasattr(fvdb_reality_capture.functional, symbol) + assert not hasattr(fvdb.functional, symbol) -def test_gaussian_splat_enums_are_owned_by_reality_capture_with_preserved_values(): + +def test_gaussian_splat_enums_are_shared_with_fvdb_with_preserved_values(): import fvdb import fvdb.viz import fvdb_reality_capture - from fvdb_reality_capture import enums - - public_enums = ( - "RollingShutterType", - "CameraModel", - "ProjectionMethod", - ) - for enum_name in public_enums: - assert getattr(fvdb_reality_capture, enum_name) is getattr(enums, enum_name) - - # Core also exposes camera and shutter enums for its functional Gaussian API. + # The camera enums are owned by fvdb-core and re-exported here as the same objects, so + # values round-trip between the two packages without conversion. for enum_name in ("RollingShutterType", "CameraModel"): - assert {member.name: member.value for member in getattr(fvdb_reality_capture, enum_name)} == { - member.name: member.value for member in getattr(fvdb, enum_name) - } + assert getattr(fvdb_reality_capture, enum_name) is getattr(fvdb, enum_name) + + # Pipeline enums are owned here; fvdb kernels never take them. + for enum_name in ("GaussianRenderMode", "ProjectionMethod"): + assert not hasattr(fvdb, enum_name) + assert {member.name: member.value for member in fvdb_reality_capture.GaussianRenderMode} == { + "FEATURES": 0, + "DEPTH": 1, + "FEATURES_AND_DEPTH": 2, + } assert not hasattr(fvdb_reality_capture, "ShOrderingMode")