mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
triangle splatting with per pixel sorting
This commit is contained in:
@@ -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]
|
||||
);
|
||||
@@ -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(
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user