From 4d04280cb55026a9be44b9dcf4ef40a755c6d28c Mon Sep 17 00:00:00 2001 From: Harry Chen Date: Tue, 4 Nov 2025 20:53:01 -0500 Subject: [PATCH] triangle splatting with per pixel sorting --- setup.py | 5 +- .../splat/cuda/_wrapper_rasterize.py | 12 + .../cuda/csrc/PrimitiveOpaqueTriangle.cuh | 6 + .../cuda/csrc/RasterizationSortedEval3DBwd.cu | 507 ++++++++++++++++++ .../csrc/RasterizationSortedEval3DBwd.cuh | 107 ++++ .../cuda/csrc/RasterizationSortedEval3DFwd.cu | 438 +++++++++++++++ .../csrc/RasterizationSortedEval3DFwd.cuh | 89 +++ spirulae_splat/splat/cuda/csrc/ext.cpp | 8 +- .../splat/cuda/csrc/generated/primitive.cu | 34 +- .../primitive_opaque_triangle_eval3d.slang | 10 + spirulae_splat/splat/rendering.py | 31 +- spirulae_splat/viewer/viewer.py | 2 +- tests/test_rasterization.py | 14 +- 13 files changed, 1208 insertions(+), 55 deletions(-) create mode 100644 spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DBwd.cu create mode 100644 spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DBwd.cuh create mode 100644 spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DFwd.cu create mode 100644 spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DFwd.cuh diff --git a/setup.py b/setup.py index 67fbc552..9b285858 100644 --- a/setup.py +++ b/setup.py @@ -24,7 +24,7 @@ def extract_function_declarations(code): (?:inline)?\s* (?:__global__|__device__)?\s* # Match the return type - (?:void|int[234]?|float[234]?|torch::Tensor|std::tuple<[\w:\s*&<>\[\],\/]+?>) + (?:void|u?int[234]?|float[234]?|torch::Tensor|std::tuple<[\w:\s*&<>\[\],\/]+?>) # Match the function name \s+\b\w+\b\s* # Match the function parameters @@ -113,6 +113,7 @@ def get_extensions(): nvcc_flags += ["-O3", "--use_fast_math"] if LINE_INFO: nvcc_flags += ["-lineinfo", "--generate-line-info", "--source-in-ptx"] + # nvcc_flags += ["-Xptxas", "-v", "-Xptxas", "--warn-on-spills"] if torch.version.hip: # USE_ROCM was added to later versions of PyTorch. # Define here to support older PyTorch versions as well: @@ -169,6 +170,8 @@ generate_header(path+"RasterizationFwd.cu", path+"RasterizationFwd.cuh") generate_header(path+"RasterizationBwd.cu", path+"RasterizationBwd.cuh") generate_header(path+"RasterizationEval3DFwd.cu", path+"RasterizationEval3DFwd.cuh") generate_header(path+"RasterizationEval3DBwd.cu", path+"RasterizationEval3DBwd.cuh") +generate_header(path+"RasterizationSortedEval3DFwd.cu", path+"RasterizationSortedEval3DFwd.cuh") +generate_header(path+"RasterizationSortedEval3DBwd.cu", path+"RasterizationSortedEval3DBwd.cuh") setup( name="spirulae_splat", diff --git a/spirulae_splat/splat/cuda/_wrapper_rasterize.py b/spirulae_splat/splat/cuda/_wrapper_rasterize.py index d4137e03..891de2a2 100644 --- a/spirulae_splat/splat/cuda/_wrapper_rasterize.py +++ b/spirulae_splat/splat/cuda/_wrapper_rasterize.py @@ -432,6 +432,9 @@ class _RasterizeToPixelsOpaqueTriangleEval3D(torch.autograd.Function): f"CameraModelType.{camera_model.upper()}" ) + # from time import perf_counter + # torch.cuda.synchronize() + # time0 = perf_counter() ( (render_rgbs, render_depths, render_normals), render_Ts, last_ids, @@ -443,6 +446,9 @@ class _RasterizeToPixelsOpaqueTriangleEval3D(torch.autograd.Function): backgrounds, masks, width, height, tile_size, isect_offsets, flatten_ids, ) + # torch.cuda.synchronize() + # time1 = perf_counter() + # print(f"fwd: {1e3*(time1-time0):.2f} ms") ctx.save_for_backward( hardness, depths, verts, rgbs, normals, @@ -486,6 +492,9 @@ class _RasterizeToPixelsOpaqueTriangleEval3D(torch.autograd.Function): height = ctx.height tile_size = ctx.tile_size + # from time import perf_counter + # torch.cuda.synchronize() + # time0 = perf_counter() ( (v_hardness, v_depths, v_verts, v_rgbs, v_normals), v_viewmats, @@ -500,6 +509,9 @@ class _RasterizeToPixelsOpaqueTriangleEval3D(torch.autograd.Function): v_render_alphas.contiguous(), (v_distortion_rgbs.contiguous(), v_distortion_depths.contiguous(), v_distortion_normals.contiguous()), ) + # torch.cuda.synchronize() + # time1 = perf_counter() + # print(f"bwd: {1e3*(time1-time0):.2f} ms") v_backgrounds = None if ctx.needs_input_grad[11]: diff --git a/spirulae_splat/splat/cuda/csrc/PrimitiveOpaqueTriangle.cuh b/spirulae_splat/splat/cuda/csrc/PrimitiveOpaqueTriangle.cuh index 51801a87..b70bb30b 100644 --- a/spirulae_splat/splat/cuda/csrc/PrimitiveOpaqueTriangle.cuh +++ b/spirulae_splat/splat/cuda/csrc/PrimitiveOpaqueTriangle.cuh @@ -791,6 +791,12 @@ struct OpaqueTriangle::WorldEval3D { return v_splat; } + __device__ __forceinline__ float evaluate_sorting_depth(float3 ray_o, float3 ray_d) { + return evaluate_sorting_depth_opaque_triangle( + &verts, &rgbs, ray_o, ray_d + ); + } + __device__ __forceinline__ OpaqueTriangle::RenderOutput evaluate_color(float3 ray_o, float3 ray_d) { float3 out_rgb; float out_depth; evaluate_color_opaque_triangle( diff --git a/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DBwd.cu b/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DBwd.cu new file mode 100644 index 00000000..1f4a5d3b --- /dev/null +++ b/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DBwd.cu @@ -0,0 +1,507 @@ +// Modified from https://github.com/nerfstudio-project/gsplat/blob/main/gsplat/cuda/csrc/RasterizeToPixels3DGSBwd.cu + +#include "RasterizationEval3DBwd.cuh" + +#include "common.cuh" + +#include +#include + +#include + + +constexpr uint BLOCK_SIZE = TILE_SIZE * TILE_SIZE; + + +template +struct MaxPriorityQueue { + uint2 arr[MAX_SIZE]; // x is key and y is sorting value + // uint size; // commented to prevent spill to local memory + + inline __device__ void push(uint key, uint& size, float value) { + uint idx = size; + arr[idx] = make_uint2(key, value <= 0.0f ? 0u : __float_as_uint(value)); + size++; + while (idx > 0) { + uint parent = (idx - 1) / 2; + if (arr[idx].y <= arr[parent].y) + break; + uint2 temp = arr[idx]; + arr[idx] = arr[parent]; + arr[parent] = temp; + idx = parent; + } + } + + inline __device__ uint2 pop(uint& size) { + uint2 result = arr[0]; + size--; + if (size > 0) { + uint idx = 0; + uint2 mval = arr[size]; + arr[0] = mval; + while (true) { + uint left = 2 * idx + 1; + uint right = 2 * idx + 2; + uint midx = idx; + if (left < size && arr[left].y > mval.y) + midx = left, mval = arr[left]; + if (right < size && arr[right].y > mval.y) + midx = right, mval = arr[right]; + if (midx == idx) + break; + uint2 temp = arr[midx]; + mval = arr[midx] = arr[idx]; + arr[idx] = temp; + idx = midx; + } + } + return result; + } + +}; + + +template +__global__ void rasterize_to_pixels_sorted_eval3d_bwd_kernel( + const uint32_t I, + const uint32_t N, + const uint32_t n_isects, + const bool packed, + // fwd inputs + typename SplatPrimitive::WorldEval3D::Buffer splat_buffer, + const float *__restrict__ viewmats, // [B, C, 4, 4] + const float *__restrict__ Ks, // [B, C, 3, 3] + const CameraDistortionCoeffsBuffer dist_coeffs, + const float *__restrict__ backgrounds, // [..., CDIM] or [nnz, CDIM] + const bool *__restrict__ masks, // [..., tile_height, tile_width] + const uint32_t image_width, + const uint32_t image_height, + const uint32_t tile_width, + const uint32_t tile_height, + const int32_t *__restrict__ tile_offsets, // [..., tile_height, tile_width] + const int32_t *__restrict__ flatten_ids, // [n_isects] + // fwd outputs + const float *__restrict__ render_Ts, // [..., image_height, image_width, 1] + const int32_t *__restrict__ last_ids, // [..., image_height, image_width] + typename SplatPrimitive::RenderOutput::Buffer render_output_buffer, + typename SplatPrimitive::RenderOutput::Buffer render2_output_buffer, + // grad outputs + typename SplatPrimitive::RenderOutput::Buffer v_render_output_buffer, + const float *__restrict__ v_render_alphas, // [..., image_height, image_width, 1] + typename SplatPrimitive::RenderOutput::Buffer v_distortions_output_buffer, + // grad inputs + typename SplatPrimitive::WorldEval3D::Buffer v_splat_buffer +) { + auto block = cg::this_thread_block(); + cg::thread_block_tile warp = cg::tiled_partition(block); + uint32_t image_id = block.group_index().x; + uint32_t tile_id = block.group_index().y * tile_width + block.group_index().z; + uint32_t thread_id = block.thread_rank(); + + tile_offsets += image_id * tile_height * tile_width; + render_Ts += image_id * image_height * image_width; + last_ids += image_id * image_height * image_width; + v_render_alphas += image_id * image_height * image_width; + if (backgrounds != nullptr) { + backgrounds += image_id; + } + if (masks != nullptr) { + masks += image_id * tile_height * tile_width; + } + + // when the mask is provided, do nothing and return if + // this tile is labeled as False + if (masks != nullptr && !masks[tile_id]) { + return; + } + + // Load camera + viewmats += image_id * 16; + Ks += image_id * 9; + float3x3 R = { + viewmats[0], viewmats[1], viewmats[2], // 1st row + viewmats[4], viewmats[5], viewmats[6], // 2nd row + viewmats[8], viewmats[9], viewmats[10], // 3rd row + }; + float3 t = { viewmats[3], viewmats[7], viewmats[11] }; + float fx = Ks[0], fy = Ks[4], cx = Ks[2], cy = Ks[5]; + float4 radial_coeffs = dist_coeffs.radial_coeffs != nullptr ? + dist_coeffs.radial_coeffs[image_id] : make_float4(0.0f); + float2 tangential_coeffs = dist_coeffs.tangential_coeffs != nullptr ? + dist_coeffs.tangential_coeffs[image_id] : make_float2(0.0f); + float2 thin_prism_coeffs = dist_coeffs.thin_prism_coeffs != nullptr ? + dist_coeffs.thin_prism_coeffs[image_id] : make_float2(0.0f); + + // arrange 16x16 tile into 8x4 subtiles (one per warp) + static_assert(TILE_SIZE == 16); + static_assert(WARP_SIZE == 32); + uint tile_idx = block.thread_index().y * TILE_SIZE + block.thread_index().x; + uint warp_idx = tile_idx / WARP_SIZE, + lane_idx = tile_idx % WARP_SIZE; + uint32_t pix_y = block.group_index().y * TILE_SIZE + (warp_idx / 2) * 4 + lane_idx / 8; + uint32_t pix_x = block.group_index().z * TILE_SIZE + (warp_idx % 2) * 8 + lane_idx % 8; + + int32_t pix_id_global = pix_y * image_width + pix_x; + uint pix_id_image_global = image_id * image_height * image_width + pix_id_global; + bool inside = (pix_x < image_width && pix_y < image_height); + + float2 pix_Ts_with_grad = { + (inside ? render_Ts[pix_id_global] : 0.0f), + (inside ? -v_render_alphas[pix_id_global] : 0.0f) + }; + typename SplatPrimitive::RenderOutput v_pix_colors = (inside ? + SplatPrimitive::RenderOutput::load(v_render_output_buffer, pix_id_image_global) + : SplatPrimitive::RenderOutput::zero()); + // float pix_background[CDIM]; // TODO + + const float px = (float)pix_x + 0.5f; + const float py = (float)pix_y + 0.5f; + float3 ray_o, ray_d; + generate_ray( + R, t, {(px-cx)/fx, (py-cy)/fy}, camera_model == gsplat::CameraModelType::FISHEYE, + radial_coeffs, tangential_coeffs, thin_prism_coeffs, + &ray_o, &ray_d + ); + + typename SplatPrimitive::RenderOutput pix_colors; + typename SplatPrimitive::RenderOutput pix2_colors; + typename SplatPrimitive::RenderOutput v_distortion_out; + if (output_distortion) { + pix_colors = (inside ? + SplatPrimitive::RenderOutput::load(render_output_buffer, pix_id_image_global) + : SplatPrimitive::RenderOutput::zero()); + pix2_colors = (inside ? + SplatPrimitive::RenderOutput::load(render2_output_buffer, pix_id_image_global) + : SplatPrimitive::RenderOutput::zero()); + v_distortion_out = (inside ? + SplatPrimitive::RenderOutput::load(v_distortions_output_buffer, pix_id_image_global) + : SplatPrimitive::RenderOutput::zero()); + } + + int32_t range_start = tile_offsets[tile_id]; + int32_t range_end = + (image_id == I - 1) && (tile_id == tile_width * tile_height - 1) + ? n_isects + : tile_offsets[tile_id + 1]; + int32_t bin_final = (inside ? last_ids[pix_id_global] : 0x7fffffff); + + static constexpr uint MAX_PQUEUE_SIZE = 32; + MaxPriorityQueue pqueue; + uint pqueue_size = 0; + + __shared__ typename SplatPrimitive::WorldEval3D splat_batch[BLOCK_SIZE]; + + for (int32_t t = range_end-1; t >= (int)range_start - (int)MAX_PQUEUE_SIZE-1; t--) { + + // load splats into shared memory + if (((range_end-1) - t) % BLOCK_SIZE == 0) { + block.sync(); + int32_t t1 = t - (int)block.thread_rank(); + if (t1 >= range_start) { + uint32_t splat_idx = flatten_ids[t1]; + typename SplatPrimitive::WorldEval3D splat = + SplatPrimitive::WorldEval3D::loadWithPrecompute(splat_buffer, splat_idx); + splat_batch[block.thread_rank()] = splat; + } + block.sync(); + } + + // early skip if done + bool active = inside && (t <= bin_final); + active &= (t >= range_start || pqueue_size != 0); + if (__ballot_sync(~0u, active) == 0) + continue; + + bool hasSplat = false; + float depth; + if (active && t >= range_start) { + // uint32_t splat_idx = flatten_ids[t]; + // typename SplatPrimitive::WorldEval3D splat = + // SplatPrimitive::WorldEval3D::loadWithPrecompute(splat_buffer, splat_idx); + typename SplatPrimitive::WorldEval3D splat = + splat_batch[((range_end-1) - t) % BLOCK_SIZE]; + float alpha = splat.evaluate_alpha(ray_o, ray_d); + hasSplat |= (alpha >= ALPHA_THRESHOLD); + if (hasSplat) + depth = splat.evaluate_sorting_depth(ray_o, ray_d); + } + + if (pqueue_size >= MAX_PQUEUE_SIZE || (t < range_start && pqueue_size != 0)) { + uint32_t t = pqueue.pop(pqueue_size).x; + uint32_t splat_idx = flatten_ids[t]; + typename SplatPrimitive::WorldEval3D splat = + SplatPrimitive::WorldEval3D::loadWithPrecompute(splat_buffer, splat_idx); + float alpha = splat.evaluate_alpha(ray_o, ray_d); + + // forward: + // \left(c_{1},T_{1}\right)=\left(c_{0}+\alpha_{i}T_{0}c_{i},\ T_{0}\left(1-\alpha_{i}\right)\right) + float T1 = pix_Ts_with_grad.x; + float v_T1 = pix_Ts_with_grad.y; + + // undo pixel: + // T_{0}=\frac{T_{1}}{1-\alpha_{i}} + float ra = 1.0f / (1.0f - alpha); + float T0 = T1 * ra; + + typename SplatPrimitive::RenderOutput color = splat.evaluate_color(ray_o, ray_d); + typename SplatPrimitive::RenderOutput v_c = v_pix_colors; + + // gradient to alpha: + // \frac{dL}{d\alpha_{i}} + // = \frac{dL}{dc_{1}}\frac{dc_{1}}{d\alpha_{i}}+\frac{dL}{dT_{1}}\frac{dT_{1}}{d\alpha_{i}} + // = T_{0}\frac{dL}{dc_{1}}c_{i}-\frac{dL}{dT_{1}}T_{0} + float v_alpha = T0 * color.dot(v_c) -v_T1 * T0; + + // gradient to color: + // \frac{dL}{dc_{i}} + // = \frac{dL}{dc_{1}}\frac{dc_{1}}{dc_{i}} + // = \alpha_{i}T_{0}\frac{dL}{dc_{1}} + typename SplatPrimitive::RenderOutput v_color = v_c * (alpha * T0); + + // update pixel gradient: + // \frac{dL}{dT_{0}} + // = \frac{dL}{dc_{1}}\frac{dc_{1}}{dT_{0}}+\frac{dL}{dT_{1}}\frac{dT_{1}}{dT_{0}} + // = \alpha_{i}\frac{dL}{dc_{1}}c_{i}+\frac{dL}{dT_{1}}\left(1-\alpha_{i}\right) + float v_T0 = alpha * color.dot(v_c) + v_T1 * (1.0f - alpha); + + // distortion + if (output_distortion) { + // \left(d_{1},s_{1}\right)=\left(d_{0}+\alpha_{i}T_{0}\left(c_{i}^{2}\left(1-T_{0}\right)-2c_{i}c_{0}+s_{0}\right),\ s_{0}+\alpha_{i}T_{0}c_{i}^{2}\right) + // \frac{dL}{ds}=0 + typename SplatPrimitive::RenderOutput v_dist = v_distortion_out; + typename SplatPrimitive::RenderOutput c0 = + pix_colors + color * -alpha * T0; + typename SplatPrimitive::RenderOutput s0 = + pix2_colors + color * color * -alpha * T0; + // \frac{dL}{d\alpha_{i}}=\frac{dL}{dd_{1}}\frac{dd_{1}}{d\alpha_{i}}=T_{0}\left(c_{i}^{2}\left(1-T_{0}\right)-2c_{i}c_{0}+s_{0}\right)\frac{dL}{dd_{1}} + v_alpha += T0 * ( + color * color * (1.0f-T0) + + color * c0 * -2.0f + s0 + ).dot(v_dist); + // \frac{dL}{dc_{i}}=\frac{dL}{dd_{1}}\frac{dd_{1}}{dc_{i}}=2\alpha_{i}T_{0}\left(c_{i}\left(1-T_{0}\right)-c_{0}\right)\frac{dL}{dd_{1}} + v_color += ( + color * (1.0f-T0) + + c0 * -1.0f + ) * v_dist * (2.0f * alpha * T0); + // \alpha_{i}\left(c_{i}^{2}\left(1-2T_{0}\right)-2c_{i}c_{0}+s_{0}\right)\frac{dL}{dd_{1}} + v_T0 += alpha * ( + color * color * (1.0f-2.0f*T0) + + color * c0 * -2.0f + s0 + ).dot(v_dist); + // undo pixel state + pix_colors = c0; + pix2_colors = s0; + } + + // backward diff splat + float3 v_ray_o, v_ray_d; // TODO + typename SplatPrimitive::WorldEval3D v_splat = SplatPrimitive::WorldEval3D::zero(); + v_splat += splat.evaluate_alpha_vjp(ray_o, ray_d, v_alpha, v_ray_o, v_ray_d); + v_splat += splat.evaluate_color_vjp(ray_o, ray_d, v_color, v_ray_o, v_ray_d); + + // update pixel states + pix_Ts_with_grad = { T0, v_T0 }; + // v_pix_colors remains the same + + // accumulate gradient + splat.atomicAddGradientToBuffer(v_splat, v_splat_buffer, splat_idx); + } + + if (hasSplat) { + pqueue.push(t, pqueue_size, depth); + } + } + // TODO: gradient to viewmat +} + + +template +inline void launch_rasterize_to_pixels_sorted_eval3d_bwd_kernel( + // Gaussian parameters + typename SplatPrimitive::WorldEval3D::Tensor splats, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., 3] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids, // [n_isects] + // forward outputs + const at::Tensor render_Ts, // [..., image_height, image_width, 1] + const at::Tensor last_ids, // [..., image_height, image_width] + typename SplatPrimitive::RenderOutput::Tensor *render_outputs, + typename SplatPrimitive::RenderOutput::Tensor *render2_outputs, + // gradients of outputs + typename SplatPrimitive::RenderOutput::Tensor v_render_outputs, + const at::Tensor v_render_alphas, // [..., image_height, image_width, 1] + typename SplatPrimitive::RenderOutput::Tensor *v_distortion_outputs, + // outputs + typename SplatPrimitive::WorldEval3D::Tensor v_splats +) { + bool packed = splats.isPacked(); + uint32_t N = packed ? 0 : splats.size(); // number of gaussians + uint32_t I = render_Ts.numel() / (image_height * image_width); // number of images + uint32_t tile_height = tile_offsets.size(-2); + uint32_t tile_width = tile_offsets.size(-1); + uint32_t n_isects = flatten_ids.size(0); + + // Each block covers a tile on the image. In total there are + // I * tile_height * tile_width blocks. + dim3 threads = {TILE_SIZE, TILE_SIZE, 1}; + dim3 grid = {I, tile_height, tile_width}; + + if (n_isects == 0) { + // skip the kernel launch if there are no elements + return; + } + + #define _LAUNCH_ARGS <<>>( \ + I, N, n_isects, packed, \ + splats.buffer(), \ + viewmats.data_ptr(), Ks.data_ptr(), dist_coeffs, \ + backgrounds.has_value() ? backgrounds.value().data_ptr() : nullptr, \ + masks.has_value() ? masks.value().data_ptr() : nullptr, \ + image_width, image_height, tile_width, tile_height, \ + tile_offsets.data_ptr(), flatten_ids.data_ptr(), \ + render_Ts.data_ptr(), last_ids.data_ptr(), \ + output_distortion ? render_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \ + output_distortion ? render2_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \ + v_render_outputs.buffer(), v_render_alphas.data_ptr(), \ + output_distortion ? v_distortion_outputs->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \ + v_splats.buffer() \ + ) + + if (camera_model == gsplat::CameraModelType::PINHOLE) + rasterize_to_pixels_sorted_eval3d_bwd_kernel _LAUNCH_ARGS; + else if (camera_model == gsplat::CameraModelType::FISHEYE) + rasterize_to_pixels_sorted_eval3d_bwd_kernel _LAUNCH_ARGS; + else + throw std::runtime_error("Unsupported camera model"); + CHECK_DEVICE_ERROR(cudaGetLastError()); + + #undef _LAUNCH_ARGS +} + + +template +inline std::tuple< + typename SplatPrimitive::WorldEval3D::TensorTuple, + std::optional // v_viewmats +> rasterize_to_pixels_sorted_eval3d_bwd_tensor( + // Gaussian parameters + typename SplatPrimitive::WorldEval3D::TensorTuple splats_tuple, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + const uint32_t tile_size, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids, // [n_isects] + // forward outputs + const at::Tensor render_Ts, // [..., image_height, image_width, 1] + const at::Tensor last_ids, // [..., image_height, image_width] + std::optional render_outputs_tuple, + std::optional render2_outputs_tuple, + // gradients of outputs + typename SplatPrimitive::RenderOutput::TensorTuple v_render_outputs, + const at::Tensor v_render_alphas, // [..., image_height, image_width, 1] + std::optional v_distortion_outputs_tuple +) { + DEVICE_GUARD(tile_offsets); + CHECK_INPUT(tile_offsets); + CHECK_INPUT(flatten_ids); + CHECK_INPUT(viewmats); + CHECK_INPUT(Ks); + CHECK_INPUT(render_Ts); + CHECK_INPUT(last_ids); + CHECK_INPUT(v_render_alphas); + if (backgrounds.has_value()) + CHECK_INPUT(backgrounds.value()); + if (masks.has_value()) + CHECK_INPUT(masks.value()); + + if (tile_size != TILE_SIZE) + AT_ERROR("Unsupported tile size"); + + typename SplatPrimitive::WorldEval3D::Tensor splats(splats_tuple); + typename SplatPrimitive::WorldEval3D::Tensor v_splats = splats.zeros_like(); + + at::Tensor v_viewmats = at::zeros_like(viewmats); // TODO + + std::optional render_outputs = std::nullopt; + std::optional render2_outputs = std::nullopt; + std::optional v_distortion_outputs = std::nullopt; + if (output_distortion) { + render_outputs = render_outputs_tuple; + render2_outputs = render2_outputs_tuple; + v_distortion_outputs = v_distortion_outputs_tuple; + } + + launch_rasterize_to_pixels_sorted_eval3d_bwd_kernel( + splats, + viewmats, Ks, camera_model, dist_coeffs, + backgrounds, masks, + image_width, image_height, tile_offsets, flatten_ids, + render_Ts, last_ids, + output_distortion ? &render_outputs.value() : nullptr, + output_distortion ? &render2_outputs.value() : nullptr, + v_render_outputs, v_render_alphas, + output_distortion ? &v_distortion_outputs.value() : nullptr, + v_splats + ); + + return std::make_tuple(v_splats.tupleAll(), v_viewmats); +} + +std::tuple< + OpaqueTriangle::WorldEval3D::TensorTuple, + std::optional // v_viewmats +> rasterize_to_pixels_opaque_triangle_sorted_eval3d_bwd( + // Gaussian parameters + OpaqueTriangle::WorldEval3D::TensorTuple splats_tuple, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + const uint32_t tile_size, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids, // [n_isects] + // forward outputs + const at::Tensor render_Ts, // [..., image_height, image_width, 1] + const at::Tensor last_ids, // [..., image_height, image_width] + std::optional render_outputs, + std::optional render2_outputs, + // gradients of outputs + OpaqueTriangle::RenderOutput::TensorTuple v_render_outputs, + const at::Tensor v_render_alphas, // [..., image_height, image_width, 1] + std::optional v_distortion_outputs +) { + return rasterize_to_pixels_sorted_eval3d_bwd_tensor( + splats_tuple, + viewmats, Ks, camera_model, dist_coeffs, + backgrounds, masks, + image_width, image_height, tile_size, tile_offsets, flatten_ids, + render_Ts, last_ids, render_outputs, render2_outputs, + v_render_outputs, v_render_alphas, v_distortion_outputs + ); +} diff --git a/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DBwd.cuh b/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DBwd.cuh new file mode 100644 index 00000000..2f72098f --- /dev/null +++ b/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DBwd.cuh @@ -0,0 +1,107 @@ +#include +#include +#include + +#include + +#include "Primitive3DGS.cuh" +#include "PrimitiveOpaqueTriangle.cuh" + +#include "PixelWise.cuh" // CameraDistortionCoeffsBuffer + + +/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */ + + + +template +inline void launch_rasterize_to_pixels_sorted_eval3d_bwd_kernel( + // Gaussian parameters + typename SplatPrimitive::WorldEval3D::Tensor splats, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., 3] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids, // [n_isects] + // forward outputs + const at::Tensor render_Ts, // [..., image_height, image_width, 1] + const at::Tensor last_ids, // [..., image_height, image_width] + typename SplatPrimitive::RenderOutput::Tensor *render_outputs, + typename SplatPrimitive::RenderOutput::Tensor *render2_outputs, + // gradients of outputs + typename SplatPrimitive::RenderOutput::Tensor v_render_outputs, + const at::Tensor v_render_alphas, // [..., image_height, image_width, 1] + typename SplatPrimitive::RenderOutput::Tensor *v_distortion_outputs, + // outputs + typename SplatPrimitive::WorldEval3D::Tensor v_splats +); + + +template +inline std::tuple< + typename SplatPrimitive::WorldEval3D::TensorTuple, + std::optional // v_viewmats +> rasterize_to_pixels_sorted_eval3d_bwd_tensor( + // Gaussian parameters + typename SplatPrimitive::WorldEval3D::TensorTuple splats_tuple, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + const uint32_t tile_size, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids, // [n_isects] + // forward outputs + const at::Tensor render_Ts, // [..., image_height, image_width, 1] + const at::Tensor last_ids, // [..., image_height, image_width] + std::optional render_outputs_tuple, + std::optional render2_outputs_tuple, + // gradients of outputs + typename SplatPrimitive::RenderOutput::TensorTuple v_render_outputs, + const at::Tensor v_render_alphas, // [..., image_height, image_width, 1] + std::optional v_distortion_outputs_tuple +); + + +std::tuple< + OpaqueTriangle::WorldEval3D::TensorTuple, + std::optional // v_viewmats +> rasterize_to_pixels_opaque_triangle_sorted_eval3d_bwd( + // Gaussian parameters + OpaqueTriangle::WorldEval3D::TensorTuple splats_tuple, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + const uint32_t tile_size, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids, // [n_isects] + // forward outputs + const at::Tensor render_Ts, // [..., image_height, image_width, 1] + const at::Tensor last_ids, // [..., image_height, image_width] + std::optional render_outputs, + std::optional render2_outputs, + // gradients of outputs + OpaqueTriangle::RenderOutput::TensorTuple v_render_outputs, + const at::Tensor v_render_alphas, // [..., image_height, image_width, 1] + std::optional v_distortion_outputs +); diff --git a/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DFwd.cu b/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DFwd.cu new file mode 100644 index 00000000..5a784721 --- /dev/null +++ b/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DFwd.cu @@ -0,0 +1,438 @@ +// Modified from https://github.com/nerfstudio-project/gsplat/blob/main/gsplat/cuda/csrc/RasterizeToPixels3DGSFwd.cu + +#include "RasterizationEval3DFwd.cuh" + +#include "common.cuh" + +#include +#include + + +constexpr uint BLOCK_SIZE = TILE_SIZE * TILE_SIZE; + + +template +struct MinPriorityQueue { + uint2 arr[MAX_SIZE]; // x is key and y is sorting value + // uint size; // commented to prevent spill to local memory + + inline __device__ void push(uint key, uint& size, float value) { + uint idx = size; + arr[idx] = make_uint2(key, value <= 0.0f ? 0u : __float_as_uint(value)); + size++; + while (idx > 0) { + uint parent = (idx - 1) / 2; + if (arr[idx].y >= arr[parent].y) + break; + uint2 temp = arr[idx]; + arr[idx] = arr[parent]; + arr[parent] = temp; + idx = parent; + } + } + + inline __device__ uint2 pop(uint& size) { + uint2 result = arr[0]; + size--; + if (size > 0) { + uint idx = 0; + uint2 mval = arr[size]; + arr[0] = mval; + while (true) { + uint left = 2 * idx + 1; + uint right = 2 * idx + 2; + uint midx = idx; + if (left < size && arr[left].y < mval.y) + midx = left, mval = arr[left]; + if (right < size && arr[right].y < mval.y) + midx = right, mval = arr[right]; + if (midx == idx) + break; + uint2 temp = arr[midx]; + mval = arr[midx] = arr[idx]; + arr[idx] = temp; + idx = midx; + } + } + return result; + } + +}; + + +template +__global__ void rasterize_to_pixels_sorted_eval3d_fwd_kernel( + const uint32_t I, + const uint32_t N, + const uint32_t n_isects, + const bool packed, + const typename SplatPrimitive::WorldEval3D::Buffer splat_buffer, + const float *__restrict__ viewmats, // [B, C, 4, 4] + const float *__restrict__ Ks, // [B, C, 3, 3] + const CameraDistortionCoeffsBuffer dist_coeffs, + const float3 *__restrict__ backgrounds, // [I, 3] + const bool *__restrict__ masks, // [I, tile_height, tile_width] + const uint32_t image_width, + const uint32_t image_height, + const uint32_t tile_width, + const uint32_t tile_height, + const int32_t *__restrict__ tile_offsets, // [I, tile_height, tile_width] + const int32_t *__restrict__ flatten_ids, // [n_isects] + typename SplatPrimitive::RenderOutput::Buffer render_colors, // [I, image_height, image_width, ...] + float *__restrict__ render_Ts, // [I, image_height, image_width, 1] + int32_t *__restrict__ last_ids, // [I, image_height, image_width] + typename SplatPrimitive::RenderOutput::Buffer render_colors2, // [I, image_height, image_width, ...] + typename SplatPrimitive::RenderOutput::Buffer render_distortions // [I, image_height, image_width, ...] +) { + // each thread draws one pixel, but also timeshares caching gaussians in a + // shared tile + + auto block = cg::this_thread_block(); + int32_t image_id = block.group_index().x; + int32_t tile_id = + block.group_index().y * tile_width + block.group_index().z; + + tile_offsets += image_id * tile_height * tile_width; + render_Ts += image_id * image_height * image_width; + last_ids += image_id * image_height * image_width; + if (backgrounds != nullptr) { + backgrounds += image_id; + } + if (masks != nullptr) { + masks += image_id * tile_height * tile_width; + } + + // arrange 16x16 tile into 8x4 subtiles (one per warp) + static_assert(TILE_SIZE == 16); + static_assert(WARP_SIZE == 32); + uint tile_idx = block.thread_index().y * TILE_SIZE + block.thread_index().x; + uint warp_idx = tile_idx / WARP_SIZE, + lane_idx = tile_idx % WARP_SIZE; + uint32_t i = block.group_index().y * TILE_SIZE + (warp_idx / 2) * 4 + lane_idx / 8; + uint32_t j = block.group_index().z * TILE_SIZE + (warp_idx % 2) * 8 + lane_idx % 8; + + float px = (float)j + 0.5f; + float py = (float)i + 0.5f; + int32_t pix_id = i * image_width + j; + + // Load camera + viewmats += image_id * 16; + Ks += image_id * 9; + float3x3 R = { + viewmats[0], viewmats[1], viewmats[2], // 1st row + viewmats[4], viewmats[5], viewmats[6], // 2nd row + viewmats[8], viewmats[9], viewmats[10], // 3rd row + }; + float3 t = { viewmats[3], viewmats[7], viewmats[11] }; + float fx = Ks[0], fy = Ks[4], cx = Ks[2], cy = Ks[5]; + float4 radial_coeffs = dist_coeffs.radial_coeffs != nullptr ? + dist_coeffs.radial_coeffs[image_id] : make_float4(0.0f); + float2 tangential_coeffs = dist_coeffs.tangential_coeffs != nullptr ? + dist_coeffs.tangential_coeffs[image_id] : make_float2(0.0f); + float2 thin_prism_coeffs = dist_coeffs.thin_prism_coeffs != nullptr ? + dist_coeffs.thin_prism_coeffs[image_id] : make_float2(0.0f); + + float3 ray_o; float3 ray_d; + generate_ray( + R, t, {(px-cx)/fx, (py-cy)/fy}, camera_model == gsplat::CameraModelType::FISHEYE, + radial_coeffs, tangential_coeffs, thin_prism_coeffs, + &ray_o, &ray_d + ); + + // return if out of bounds + // keep not rasterizing threads around for reading data + bool inside = (i < image_height && j < image_width); + bool done = !inside; + + // when the mask is provided, render the background color and return + // if this tile is labeled as False + if (masks != nullptr && inside && !masks[tile_id]) { + // TODO + // render_colors[pix_id] = backgrounds == nullptr ? + // SplatPrimitive::RenderOutput(make_float3(0.f)) : + // SplatPrimitive::RenderOutput(*backgrounds); + return; + } + + // have all threads in tile process the same gaussians in batches + // first collect gaussians between range.x and range.y in batches + // which gaussians to look through in this tile + int32_t range_start = tile_offsets[tile_id]; + int32_t range_end = + (image_id == I - 1) && (tile_id == tile_width * tile_height - 1) + ? n_isects + : tile_offsets[tile_id + 1]; + + // current visibility left to render + // transmittance is gonna be used in the backward pass which requires a high + // numerical precision so we use double for it. However double make bwd 1.5x + // slower so we stick with float for now. + float T = 1.0f; + // index of most recent gaussian to write to this thread's pixel + uint32_t cur_idx = 0; + + typename SplatPrimitive::RenderOutput pix_out = SplatPrimitive::RenderOutput::zero(); + typename SplatPrimitive::RenderOutput pix2_out = SplatPrimitive::RenderOutput::zero(); + typename SplatPrimitive::RenderOutput distortion_out = SplatPrimitive::RenderOutput::zero(); + + static constexpr uint MAX_PQUEUE_SIZE = 32; + MinPriorityQueue pqueue; + uint pqueue_size = 0; + + __shared__ typename SplatPrimitive::WorldEval3D splat_batch[BLOCK_SIZE]; + + for (uint32_t t = range_start; t < range_end + MAX_PQUEUE_SIZE + 1; ++t) { + + // load splats into shared memory + if ((t - range_start) % BLOCK_SIZE == 0) { + block.sync(); + int32_t t1 = t + (int)block.thread_rank(); + if (t1 < range_end) { + uint32_t splat_idx = flatten_ids[t1]; + typename SplatPrimitive::WorldEval3D splat = + SplatPrimitive::WorldEval3D::loadWithPrecompute(splat_buffer, splat_idx); + splat_batch[block.thread_rank()] = splat; + } + block.sync(); + } + + // early skip if done + done |= (t >= range_end && pqueue_size == 0); + if (__ballot_sync(~0u, !done) == 0) + break; + + bool hasSplat = false; + float depth; + if (!done && t < range_end) { + // uint32_t splat_idx = flatten_ids[t]; + // typename SplatPrimitive::WorldEval3D splat = + // SplatPrimitive::WorldEval3D::loadWithPrecompute(splat_buffer, splat_idx); + typename SplatPrimitive::WorldEval3D splat = + splat_batch[(t - range_start) % BLOCK_SIZE]; + float alpha = splat.evaluate_alpha(ray_o, ray_d); + hasSplat |= (alpha >= ALPHA_THRESHOLD); + if (hasSplat) + depth = splat.evaluate_sorting_depth(ray_o, ray_d); + } + + if (pqueue_size >= MAX_PQUEUE_SIZE || (t >= range_end && pqueue_size != 0)) { + uint32_t t = pqueue.pop(pqueue_size).x; + uint32_t splat_idx = flatten_ids[t]; + typename SplatPrimitive::WorldEval3D splat = + SplatPrimitive::WorldEval3D::loadWithPrecompute(splat_buffer, splat_idx); + float alpha = splat.evaluate_alpha(ray_o, ray_d); + + const float next_T = T * (1.0f - alpha); + if (next_T <= 1e-4f && false) { // this pixel is done: exclusive + done = true; + pqueue_size = 0; + } + else { + const float vis = alpha * T; + const typename SplatPrimitive::RenderOutput color = splat.evaluate_color(ray_o, ray_d); + + if (output_distortion) { + distortion_out += ( + color * color * (1.0f - T) + + color * pix_out * -2.0f + + pix2_out + ) * vis; + pix2_out += color * color * vis; + } + pix_out += color * vis; + cur_idx = t; + + T = next_T; + } + } + + if (hasSplat) { + pqueue.push(t, pqueue_size, depth); + } + } + + if (inside) { + render_Ts[pix_id] = T; + int pix_id_global = image_id * image_height * image_width + pix_id; + // TODO: blend background + pix_out.saveParamsToBuffer(render_colors, pix_id_global); + // index in bin of last gaussian in this pixel + last_ids[pix_id] = static_cast(cur_idx); + // distortion + if (output_distortion) { + pix2_out.saveParamsToBuffer(render_colors2, pix_id_global); + distortion_out.saveParamsToBuffer(render_distortions, pix_id_global); + } + } +} + +template +inline void launch_rasterize_to_pixels_sorted_eval3d_fwd_kernel( + // Gaussian parameters + typename SplatPrimitive::WorldEval3D::Tensor splats, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids, // [n_isects] + // outputs + typename SplatPrimitive::RenderOutput::Tensor renders, + at::Tensor transmittances, // [..., image_height, image_width] + at::Tensor last_ids, // [..., image_height, image_width] + typename SplatPrimitive::RenderOutput::Tensor *renders2, + typename SplatPrimitive::RenderOutput::Tensor *distortions +) { + bool packed = splats.isPacked(); + uint32_t N = packed ? 0 : splats.size(); // number of gaussians + uint32_t I = transmittances.numel() / (image_height * image_width); // number of images + uint32_t tile_height = tile_offsets.size(-2); + uint32_t tile_width = tile_offsets.size(-1); + uint32_t n_isects = flatten_ids.size(0); + + // Each block covers a tile on the image. In total there are + // I * tile_height * tile_width blocks. + dim3 threads = {TILE_SIZE, TILE_SIZE, 1}; + dim3 grid = {I, tile_height, tile_width}; + + #define _LAUNCH_ARGS <<>>( \ + I, N, n_isects, packed, \ + splats.buffer(), \ + viewmats.data_ptr(), Ks.data_ptr(), dist_coeffs, \ + backgrounds.has_value() ? (float3*)backgrounds.value().data_ptr() : nullptr, \ + masks.has_value() ? masks.value().data_ptr() : nullptr, \ + image_width, image_height, tile_width, tile_height, \ + tile_offsets.data_ptr(), flatten_ids.data_ptr(), \ + renders, transmittances.data_ptr(), last_ids.data_ptr(), \ + output_distortion ? renders2->buffer() : typename SplatPrimitive::RenderOutput::Buffer(), \ + output_distortion ? distortions->buffer() : typename SplatPrimitive::RenderOutput::Buffer() \ + ) + + if (camera_model == gsplat::CameraModelType::PINHOLE) + rasterize_to_pixels_sorted_eval3d_fwd_kernel _LAUNCH_ARGS; + else if (camera_model == gsplat::CameraModelType::FISHEYE) + rasterize_to_pixels_sorted_eval3d_fwd_kernel _LAUNCH_ARGS; + else + throw std::runtime_error("Unsupported camera model"); + + #undef _LAUNCH_ARGS +} + + +template +inline std::tuple< + typename SplatPrimitive::RenderOutput::TensorTuple, + at::Tensor, + at::Tensor, + std::optional, + std::optional +> rasterize_to_pixels_sorted_eval3d_fwd_tensor( + // Gaussian parameters + typename SplatPrimitive::WorldEval3D::TensorTuple splats_tuple, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids // [n_isects] +) { + DEVICE_GUARD(tile_offsets); + CHECK_INPUT(tile_offsets); + CHECK_INPUT(flatten_ids); + CHECK_INPUT(viewmats); + CHECK_INPUT(Ks); + if (backgrounds.has_value()) + CHECK_INPUT(backgrounds.value()); + if (masks.has_value()) + CHECK_INPUT(masks.value()); + + typename SplatPrimitive::WorldEval3D::Tensor splats(splats_tuple); + + auto opt = splats.options(); + at::DimVector image_dims(tile_offsets.sizes().slice(0, tile_offsets.dim() - 2)); + + at::DimVector renders_dims(image_dims); + renders_dims.append({image_height, image_width}); + typename SplatPrimitive::RenderOutput::Tensor renders = + SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt); + + std::optional renders2 = std::nullopt; + std::optional distortions = std::nullopt; + if (output_distortion) { + renders2 = SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt); + distortions = SplatPrimitive::RenderOutput::Tensor::empty(renders_dims, opt); + } + + at::DimVector transmittance_dims(image_dims); + transmittance_dims.append({image_height, image_width, 1}); + at::Tensor transmittances = at::empty(transmittance_dims, opt); + + at::DimVector last_ids_dims(image_dims); + last_ids_dims.append({image_height, image_width}); + at::Tensor last_ids = at::empty(last_ids_dims, opt.dtype(at::kInt)); + + launch_rasterize_to_pixels_sorted_eval3d_fwd_kernel( + splats, + viewmats, Ks, camera_model, dist_coeffs, + backgrounds, masks, + image_width, image_height, tile_offsets, flatten_ids, + renders, transmittances, last_ids, + output_distortion ? &renders2.value() : nullptr, + output_distortion ? &distortions.value() : nullptr + ); + + if (output_distortion) + return std::make_tuple(renders.tuple(), transmittances, last_ids, + renders2.value().tuple(), distortions.value().tuple()); + return std::make_tuple(renders.tuple(), transmittances, last_ids, + std::nullopt, std::nullopt); +} + + +std::tuple< + OpaqueTriangle::RenderOutput::TensorTuple, + at::Tensor, + at::Tensor, + std::optional, + std::optional +> rasterize_to_pixels_opaque_triangle_sorted_eval3d_fwd( + // Gaussian parameters + OpaqueTriangle::WorldEval3D::TensorTuple splats_tuple, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + const uint32_t tile_size, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids // [n_isects] +) { + if (tile_size != TILE_SIZE) + AT_ERROR("Tile size must be " + std::to_string(TILE_SIZE)); + return rasterize_to_pixels_sorted_eval3d_fwd_tensor( + splats_tuple, + viewmats, Ks, camera_model, dist_coeffs, + backgrounds, masks, + image_width, image_height, + tile_offsets, flatten_ids + ); +} diff --git a/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DFwd.cuh b/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DFwd.cuh new file mode 100644 index 00000000..a2dc0f6c --- /dev/null +++ b/spirulae_splat/splat/cuda/csrc/RasterizationSortedEval3DFwd.cuh @@ -0,0 +1,89 @@ +#include +#include +#include + +#include + +#include "Primitive3DGS.cuh" +#include "PrimitiveOpaqueTriangle.cuh" + +#include "PixelWise.cuh" // CameraDistortionCoeffsBuffer + + +/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */ + + + +template +inline void launch_rasterize_to_pixels_sorted_eval3d_fwd_kernel( + // Gaussian parameters + typename SplatPrimitive::WorldEval3D::Tensor splats, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids, // [n_isects] + // outputs + typename SplatPrimitive::RenderOutput::Tensor renders, + at::Tensor transmittances, // [..., image_height, image_width] + at::Tensor last_ids, // [..., image_height, image_width] + typename SplatPrimitive::RenderOutput::Tensor *renders2, + typename SplatPrimitive::RenderOutput::Tensor *distortions +); + + +template +inline std::tuple< + typename SplatPrimitive::RenderOutput::TensorTuple, + at::Tensor, + at::Tensor, + std::optional, + std::optional +> rasterize_to_pixels_sorted_eval3d_fwd_tensor( + // Gaussian parameters + typename SplatPrimitive::WorldEval3D::TensorTuple splats_tuple, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids // [n_isects] +); + + +std::tuple< + OpaqueTriangle::RenderOutput::TensorTuple, + at::Tensor, + at::Tensor, + std::optional, + std::optional +> rasterize_to_pixels_opaque_triangle_sorted_eval3d_fwd( + // Gaussian parameters + OpaqueTriangle::WorldEval3D::TensorTuple splats_tuple, + const at::Tensor viewmats, // [..., C, 4, 4] + const at::Tensor Ks, // [..., C, 3, 3] + const gsplat::CameraModelType camera_model, + const CameraDistortionCoeffsTensor dist_coeffs, + const std::optional backgrounds, // [..., channels] + const std::optional masks, // [..., tile_height, tile_width] + // image size + const uint32_t image_width, + const uint32_t image_height, + const uint32_t tile_size, + // intersections + const at::Tensor tile_offsets, // [..., tile_height, tile_width] + const at::Tensor flatten_ids // [n_isects] +); diff --git a/spirulae_splat/splat/cuda/csrc/ext.cpp b/spirulae_splat/splat/cuda/csrc/ext.cpp index 101f4381..d6629173 100644 --- a/spirulae_splat/splat/cuda/csrc/ext.cpp +++ b/spirulae_splat/splat/cuda/csrc/ext.cpp @@ -15,6 +15,8 @@ #include "RasterizationBwd.cuh" #include "RasterizationEval3DFwd.cuh" #include "RasterizationEval3DBwd.cuh" +#include "RasterizationSortedEval3DFwd.cuh" +#include "RasterizationSortedEval3DBwd.cuh" #define TORCH_INDUCTOR_CPP_WRAPPER #include @@ -84,7 +86,9 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { // RasterizationEval3DFwd.cuh and RasterizationEval3DBwd.cuh m.def("rasterization_3dgs_eval3d_forward", &rasterize_to_pixels_3dgs_eval3d_fwd); m.def("rasterization_3dgs_eval3d_backward", &rasterize_to_pixels_3dgs_eval3d_bwd); - m.def("rasterization_opaque_triangle_eval3d_forward", &rasterize_to_pixels_opaque_triangle_eval3d_fwd); - m.def("rasterization_opaque_triangle_eval3d_backward", &rasterize_to_pixels_opaque_triangle_eval3d_bwd); + // m.def("rasterization_opaque_triangle_eval3d_forward", &rasterize_to_pixels_opaque_triangle_eval3d_fwd); + // m.def("rasterization_opaque_triangle_eval3d_backward", &rasterize_to_pixels_opaque_triangle_eval3d_bwd); + m.def("rasterization_opaque_triangle_eval3d_forward", &rasterize_to_pixels_opaque_triangle_sorted_eval3d_fwd); + m.def("rasterization_opaque_triangle_eval3d_backward", &rasterize_to_pixels_opaque_triangle_sorted_eval3d_bwd); } diff --git a/spirulae_splat/splat/cuda/csrc/generated/primitive.cu b/spirulae_splat/splat/cuda/csrc/generated/primitive.cu index bb8f64b4..c155f8e1 100644 --- a/spirulae_splat/splat/cuda/csrc/generated/primitive.cu +++ b/spirulae_splat/splat/cuda/csrc/generated/primitive.cu @@ -12329,18 +12329,24 @@ inline __device__ void evaluate_alpha_opaque_triangle_vjp(FixedArray return; } -inline __device__ void evaluate_color_opaque_triangle(FixedArray * verts_6, FixedArray * rgbs_4, float3 ray_o_7, float3 ray_d_7, float3 * color_13, float * depth_20) +inline __device__ float evaluate_sorting_depth_opaque_triangle(FixedArray * verts_6, FixedArray * rgbs_4, float3 ray_o_7, float3 ray_d_7) { - float3 v1v0_2 = (*verts_6)[int(1)] - (*verts_6)[int(0)]; - float3 v2v0_2 = (*verts_6)[int(2)] - (*verts_6)[int(0)]; - float3 rov0_2 = ray_o_7 - (*verts_6)[int(0)]; - float3 n_1 = cross_0(v1v0_2, v2v0_2); - float3 q_3 = cross_0(rov0_2, ray_d_7); - float d_2 = 1.0f / dot_0(ray_d_7, n_1); + float3 n_1 = cross_0((*verts_6)[int(1)] - (*verts_6)[int(0)], (*verts_6)[int(2)] - (*verts_6)[int(0)]); + return 1.0f / dot_0(ray_d_7, n_1) * dot_0(- n_1, ray_o_7 - (*verts_6)[int(0)]); +} + +inline __device__ void evaluate_color_opaque_triangle(FixedArray * verts_7, FixedArray * rgbs_5, float3 ray_o_8, float3 ray_d_8, float3 * color_13, float * depth_20) +{ + float3 v1v0_2 = (*verts_7)[int(1)] - (*verts_7)[int(0)]; + float3 v2v0_2 = (*verts_7)[int(2)] - (*verts_7)[int(0)]; + float3 rov0_2 = ray_o_8 - (*verts_7)[int(0)]; + float3 n_2 = cross_0(v1v0_2, v2v0_2); + float3 q_3 = cross_0(rov0_2, ray_d_8); + float d_2 = 1.0f / dot_0(ray_d_8, n_2); float u_31 = d_2 * dot_0(- q_3, v2v0_2); float v_31 = d_2 * dot_0(q_3, v1v0_2); - *depth_20 = d_2 * dot_0(- n_1, rov0_2); - *color_13 = (*rgbs_4)[int(0)] * make_float3 (1.0f - u_31 - v_31) + (*rgbs_4)[int(1)] * make_float3 (u_31) + (*rgbs_4)[int(2)] * make_float3 (v_31); + *depth_20 = d_2 * dot_0(- n_2, rov0_2); + *color_13 = (*rgbs_5)[int(0)] * make_float3 (1.0f - u_31 - v_31) + (*rgbs_5)[int(1)] * make_float3 (u_31) + (*rgbs_5)[int(2)] * make_float3 (v_31); *depth_20 = (F32_log(((F32_max((*depth_20), (9.999999960041972e-13f)))))); return; } @@ -12476,21 +12482,21 @@ inline __device__ void s_bwd_evaluate_color_opaque_triangle_1(DiffPair_arrayx3Cv return; } -inline __device__ void evaluate_color_opaque_triangle_vjp(FixedArray * verts_7, FixedArray * rgbs_5, float3 ray_o_8, float3 ray_d_8, float3 v_color_1, float v_depth_11, FixedArray * v_verts_3, FixedArray * v_rgbs_2, float3 * v_ray_o_4, float3 * v_ray_d_4) +inline __device__ void evaluate_color_opaque_triangle_vjp(FixedArray * verts_8, FixedArray * rgbs_6, float3 ray_o_9, float3 ray_d_9, float3 v_color_1, float v_depth_11, FixedArray * v_verts_3, FixedArray * v_rgbs_2, float3 * v_ray_o_4, float3 * v_ray_d_4) { float3 _S4788 = make_float3 (0.0f); FixedArray _S4789 = { _S4788, _S4788, _S4788 }; DiffPair_arrayx3Cvectorx3Cfloatx2C3x3Ex2C3x3E_0 dp_verts_1; - (&dp_verts_1)->primal_0 = *verts_7; + (&dp_verts_1)->primal_0 = *verts_8; (&dp_verts_1)->differential_0 = _S4789; DiffPair_arrayx3Cvectorx3Cfloatx2C3x3Ex2C3x3E_0 dp_rgbs_0; - (&dp_rgbs_0)->primal_0 = *rgbs_5; + (&dp_rgbs_0)->primal_0 = *rgbs_6; (&dp_rgbs_0)->differential_0 = _S4789; DiffPair_vectorx3Cfloatx2C3x3E_0 dp_ray_o_3; - (&dp_ray_o_3)->primal_0 = ray_o_8; + (&dp_ray_o_3)->primal_0 = ray_o_9; (&dp_ray_o_3)->differential_0 = _S4788; DiffPair_vectorx3Cfloatx2C3x3E_0 dp_ray_d_3; - (&dp_ray_d_3)->primal_0 = ray_d_8; + (&dp_ray_d_3)->primal_0 = ray_d_9; (&dp_ray_d_3)->differential_0 = _S4788; s_bwd_evaluate_color_opaque_triangle_1(&dp_verts_1, &dp_rgbs_0, &dp_ray_o_3, &dp_ray_d_3, v_color_1, v_depth_11); *v_verts_3 = (&dp_verts_1)->differential_0; diff --git a/spirulae_splat/splat/cuda/slang/primitive_opaque_triangle_eval3d.slang b/spirulae_splat/splat/cuda/slang/primitive_opaque_triangle_eval3d.slang index 657ac299..725c2c46 100644 --- a/spirulae_splat/splat/cuda/slang/primitive_opaque_triangle_eval3d.slang +++ b/spirulae_splat/splat/cuda/slang/primitive_opaque_triangle_eval3d.slang @@ -304,6 +304,16 @@ void evaluate_alpha_opaque_triangle_vjp( v_ray_d = dp_ray_d.d; } +[CudaDeviceExport] +float evaluate_sorting_depth_opaque_triangle( + float3 verts[3], float3 rgbs[3], + float3 ray_o, float3 ray_d +) { + float u, v, t; + ray_triangle_intersection_uvt(ray_o, ray_d, verts, u, v, t); + return t; +} + [CudaDeviceExport] [Differentiable] void evaluate_color_opaque_triangle( diff --git a/spirulae_splat/splat/rendering.py b/spirulae_splat/splat/rendering.py index d613d736..1cb104f9 100644 --- a/spirulae_splat/splat/rendering.py +++ b/spirulae_splat/splat/rendering.py @@ -403,7 +403,7 @@ def rasterization( ) else: - # Project Gaussians to 2D + # Project splats to 2D proj_results = fully_fused_projection( primitive, with_eval3d, @@ -589,35 +589,6 @@ def rasterization( proj_opacities = reshape_view(C, proj_opacities, N_world) colors = reshape_view(C, colors, N_world) - # Rasterize to pixels - # if render_mode in ["RGB+D+N", "RGB+ED+N"]: - # assert normals is not None, "Primitive does not support normal" - # colors = torch.cat((colors, depths[..., None], normals), dim=-1) - # if backgrounds is not None: - # backgrounds = torch.cat( - # [ - # backgrounds, - # torch.zeros(batch_dims + (C, 1), device=backgrounds.device), - # ], - # dim=-1, - # ) - # elif render_mode in ["RGB+D", "RGB+ED"]: - # colors = torch.cat((colors, depths[..., None]), dim=-1) - # if backgrounds is not None: - # backgrounds = torch.cat( - # [ - # backgrounds, - # torch.zeros(batch_dims + (C, 1), device=backgrounds.device), - # ], - # dim=-1, - # ) - # elif render_mode in ["D", "ED"]: - # colors = depths[..., None] - # if backgrounds is not None: - # backgrounds = torch.zeros(batch_dims + (C, 1), device=backgrounds.device) - # else: # RGB - # pass - # Identify intersecting tiles tile_width = math.ceil(width / float(tile_size)) tile_height = math.ceil(height / float(tile_size)) diff --git a/spirulae_splat/viewer/viewer.py b/spirulae_splat/viewer/viewer.py index bb58ffe7..142d88af 100644 --- a/spirulae_splat/viewer/viewer.py +++ b/spirulae_splat/viewer/viewer.py @@ -241,7 +241,7 @@ if __name__ == "__main__": sys.exit(-1) # Load camera - camera_path = os.path.join(os.path.dirname(__file__), "cameras/s21.yaml") + camera_path = os.path.join(os.path.dirname(__file__), "cameras/zipnerf.yaml") camera = Camera(camera_path) # Initialize model diff --git a/tests/test_rasterization.py b/tests/test_rasterization.py index d63364bf..94924a5b 100644 --- a/tests/test_rasterization.py +++ b/tests/test_rasterization.py @@ -27,10 +27,10 @@ def rasterize_ssplat(means, quats, scales, opacities, features_dc, features_sh, camera_model = ["pinhole", "fisheye"][IS_FISHEYE] quats = torch.nn.functional.normalize(quats, dim=-1) rgbd, alpha, meta = ssplat_rasterization( - primitive=["3dgs", "mip"][IS_ANTIALIASED], - splat_params=(means, quats, scales, opacities, features_dc, features_sh), - # primitive="opaque_triangle", - # splat_params=(means, quats, scales, opacities.unsqueeze(-1).repeat(1, 2), features_dc, features_sh, features_dc.unsqueeze(-2).repeat(1, 2, 1)), + # primitive=["3dgs", "mip"][IS_ANTIALIASED], + # splat_params=(means, quats, scales, opacities, features_dc, features_sh), + primitive="opaque_triangle", + splat_params=(means, quats, scales+1.8, opacities.unsqueeze(-1).repeat(1, 2), features_dc, features_sh, features_dc.unsqueeze(-2).repeat(1, 2, 1)), viewmats=viewmats, # [C, 4, 4] Ks=Ks, # [C, 3, 3] width=W, @@ -194,9 +194,9 @@ def profile_rasterization(): if __name__ == "__main__": - N = 1000 - test_rasterization() - print() + # N = 1000 + # test_rasterization() + # print() N = 200000 profile_rasterization()