triangle splatting with per pixel sorting

This commit is contained in:
Harry Chen
2025-11-04 20:53:01 -05:00
parent eee2e890c9
commit 4d04280cb5
13 changed files with 1208 additions and 55 deletions
+4 -1
View File
@@ -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",
@@ -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]:
@@ -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(
@@ -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 <gsplat/Common.h>
#include <gsplat/Utils.cuh>
#include <cub/cub.cuh>
constexpr uint BLOCK_SIZE = TILE_SIZE * TILE_SIZE;
template<uint MAX_SIZE>
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 <typename SplatPrimitive, gsplat::CameraModelType camera_model, bool output_distortion>
__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_SIZE> warp = cg::tiled_partition<WARP_SIZE>(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<MAX_PQUEUE_SIZE> 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 <typename SplatPrimitive, bool output_distortion>
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<at::Tensor> backgrounds, // [..., 3]
const std::optional<at::Tensor> 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 <<<grid, threads>>>( \
I, N, n_isects, packed, \
splats.buffer(), \
viewmats.data_ptr<float>(), Ks.data_ptr<float>(), dist_coeffs, \
backgrounds.has_value() ? backgrounds.value().data_ptr<float>() : nullptr, \
masks.has_value() ? masks.value().data_ptr<bool>() : nullptr, \
image_width, image_height, tile_width, tile_height, \
tile_offsets.data_ptr<int32_t>(), flatten_ids.data_ptr<int32_t>(), \
render_Ts.data_ptr<float>(), last_ids.data_ptr<int32_t>(), \
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<float>(), \
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<SplatPrimitive,
gsplat::CameraModelType::PINHOLE, output_distortion> _LAUNCH_ARGS;
else if (camera_model == gsplat::CameraModelType::FISHEYE)
rasterize_to_pixels_sorted_eval3d_bwd_kernel<SplatPrimitive,
gsplat::CameraModelType::FISHEYE, output_distortion> _LAUNCH_ARGS;
else
throw std::runtime_error("Unsupported camera model");
CHECK_DEVICE_ERROR(cudaGetLastError());
#undef _LAUNCH_ARGS
}
template<typename SplatPrimitive, bool output_distortion>
inline std::tuple<
typename SplatPrimitive::WorldEval3D::TensorTuple,
std::optional<at::Tensor> // 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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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<typename SplatPrimitive::RenderOutput::TensorTuple> render_outputs_tuple,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> 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<typename SplatPrimitive::RenderOutput::TensorTuple> 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<typename SplatPrimitive::RenderOutput::Tensor> render_outputs = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> render2_outputs = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> 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<SplatPrimitive, output_distortion>(
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<at::Tensor> // 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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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<typename OpaqueTriangle::RenderOutput::TensorTuple> render_outputs,
std::optional<typename OpaqueTriangle::RenderOutput::TensorTuple> render2_outputs,
// gradients of outputs
OpaqueTriangle::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename OpaqueTriangle::RenderOutput::TensorTuple> v_distortion_outputs
) {
return rasterize_to_pixels_sorted_eval3d_bwd_tensor<OpaqueTriangle, true>(
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
);
}
@@ -0,0 +1,107 @@
#include <cuda.h>
#include <cuda_runtime.h>
#include <cstdint>
#include <torch/types.h>
#include "Primitive3DGS.cuh"
#include "PrimitiveOpaqueTriangle.cuh"
#include "PixelWise.cuh" // CameraDistortionCoeffsBuffer
/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */
template <typename SplatPrimitive, bool output_distortion>
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<at::Tensor> backgrounds, // [..., 3]
const std::optional<at::Tensor> 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<typename SplatPrimitive, bool output_distortion>
inline std::tuple<
typename SplatPrimitive::WorldEval3D::TensorTuple,
std::optional<at::Tensor> // 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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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<typename SplatPrimitive::RenderOutput::TensorTuple> render_outputs_tuple,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple> 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<typename SplatPrimitive::RenderOutput::TensorTuple> v_distortion_outputs_tuple
);
std::tuple<
OpaqueTriangle::WorldEval3D::TensorTuple,
std::optional<at::Tensor> // 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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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<typename OpaqueTriangle::RenderOutput::TensorTuple> render_outputs,
std::optional<typename OpaqueTriangle::RenderOutput::TensorTuple> render2_outputs,
// gradients of outputs
OpaqueTriangle::RenderOutput::TensorTuple v_render_outputs,
const at::Tensor v_render_alphas, // [..., image_height, image_width, 1]
std::optional<typename OpaqueTriangle::RenderOutput::TensorTuple> v_distortion_outputs
);
@@ -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 <gsplat/Common.h>
#include <gsplat/Utils.cuh>
constexpr uint BLOCK_SIZE = TILE_SIZE * TILE_SIZE;
template<uint MAX_SIZE>
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 <typename SplatPrimitive, gsplat::CameraModelType camera_model, bool output_distortion>
__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<MAX_PQUEUE_SIZE> 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<int32_t>(cur_idx);
// distortion
if (output_distortion) {
pix2_out.saveParamsToBuffer(render_colors2, pix_id_global);
distortion_out.saveParamsToBuffer(render_distortions, pix_id_global);
}
}
}
template <typename SplatPrimitive, bool output_distortion>
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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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 <<<grid, threads>>>( \
I, N, n_isects, packed, \
splats.buffer(), \
viewmats.data_ptr<float>(), Ks.data_ptr<float>(), dist_coeffs, \
backgrounds.has_value() ? (float3*)backgrounds.value().data_ptr<float>() : nullptr, \
masks.has_value() ? masks.value().data_ptr<bool>() : nullptr, \
image_width, image_height, tile_width, tile_height, \
tile_offsets.data_ptr<int32_t>(), flatten_ids.data_ptr<int32_t>(), \
renders, transmittances.data_ptr<float>(), last_ids.data_ptr<int32_t>(), \
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<SplatPrimitive,
gsplat::CameraModelType::PINHOLE, output_distortion> _LAUNCH_ARGS;
else if (camera_model == gsplat::CameraModelType::FISHEYE)
rasterize_to_pixels_sorted_eval3d_fwd_kernel<SplatPrimitive,
gsplat::CameraModelType::FISHEYE, output_distortion> _LAUNCH_ARGS;
else
throw std::runtime_error("Unsupported camera model");
#undef _LAUNCH_ARGS
}
template <typename SplatPrimitive, bool output_distortion>
inline std::tuple<
typename SplatPrimitive::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple>,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple>
> 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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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<typename SplatPrimitive::RenderOutput::Tensor> renders2 = std::nullopt;
std::optional<typename SplatPrimitive::RenderOutput::Tensor> 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<SplatPrimitive, output_distortion>(
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<OpaqueTriangle::RenderOutput::TensorTuple>,
std::optional<OpaqueTriangle::RenderOutput::TensorTuple>
> 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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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<OpaqueTriangle, true>(
splats_tuple,
viewmats, Ks, camera_model, dist_coeffs,
backgrounds, masks,
image_width, image_height,
tile_offsets, flatten_ids
);
}
@@ -0,0 +1,89 @@
#include <cuda.h>
#include <cuda_runtime.h>
#include <cstdint>
#include <torch/types.h>
#include "Primitive3DGS.cuh"
#include "PrimitiveOpaqueTriangle.cuh"
#include "PixelWise.cuh" // CameraDistortionCoeffsBuffer
/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */
template <typename SplatPrimitive, bool output_distortion>
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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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 <typename SplatPrimitive, bool output_distortion>
inline std::tuple<
typename SplatPrimitive::RenderOutput::TensorTuple,
at::Tensor,
at::Tensor,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple>,
std::optional<typename SplatPrimitive::RenderOutput::TensorTuple>
> 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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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<OpaqueTriangle::RenderOutput::TensorTuple>,
std::optional<OpaqueTriangle::RenderOutput::TensorTuple>
> 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<at::Tensor> backgrounds, // [..., channels]
const std::optional<at::Tensor> 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]
);
+6 -2
View File
@@ -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 <torch/extension.h>
@@ -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);
}
@@ -12329,18 +12329,24 @@ inline __device__ void evaluate_alpha_opaque_triangle_vjp(FixedArray<float3 , 3>
return;
}
inline __device__ void evaluate_color_opaque_triangle(FixedArray<float3 , 3> * verts_6, FixedArray<float3 , 3> * 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<float3 , 3> * verts_6, FixedArray<float3 , 3> * 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<float3 , 3> * verts_7, FixedArray<float3 , 3> * 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<float3 , 3> * verts_7, FixedArray<float3 , 3> * rgbs_5, float3 ray_o_8, float3 ray_d_8, float3 v_color_1, float v_depth_11, FixedArray<float3 , 3> * v_verts_3, FixedArray<float3 , 3> * v_rgbs_2, float3 * v_ray_o_4, float3 * v_ray_d_4)
inline __device__ void evaluate_color_opaque_triangle_vjp(FixedArray<float3 , 3> * verts_8, FixedArray<float3 , 3> * rgbs_6, float3 ray_o_9, float3 ray_d_9, float3 v_color_1, float v_depth_11, FixedArray<float3 , 3> * v_verts_3, FixedArray<float3 , 3> * v_rgbs_2, float3 * v_ray_o_4, float3 * v_ray_d_4)
{
float3 _S4788 = make_float3 (0.0f);
FixedArray<float3 , 3> _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;
@@ -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(
+1 -30
View File
@@ -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))
+1 -1
View File
@@ -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
+7 -7
View File
@@ -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()