optimize rasterization forward

This commit is contained in:
Harry Chen
2026-06-17 17:28:56 -04:00
parent 9a71caa66f
commit 89b1bc7b21
3 changed files with 318 additions and 231 deletions
+5 -6
View File
@@ -35,13 +35,12 @@ inline constexpr int WARP_SIZE = 32;
inline constexpr int TILE_SIZE_X = 8;
inline constexpr int TILE_SIZE_Y = 8;
inline constexpr int MACRO_TILE_SIZE_X = 8;
inline constexpr int MACRO_TILE_SIZE_Y = 4;
inline constexpr int MACRO_TILE_SIZE_X = 2;
inline constexpr int MACRO_TILE_SIZE_Y = 2;
static_assert(
(MACRO_TILE_SIZE_X * MACRO_TILE_SIZE_Y == 1) ||
(MACRO_TILE_SIZE_X * MACRO_TILE_SIZE_Y == WARP_SIZE)
);
// The per-microtile forward and the backward work for any macro size; only the
// (retired) macro-block forward required MACRO_NUM_TILES == WARP_SIZE.
static_assert(MACRO_TILE_SIZE_X >= 1 && MACRO_TILE_SIZE_Y >= 1);
inline constexpr float ALPHA_THRESHOLD = (1.f/255.f);
@@ -35,8 +35,13 @@ namespace SlangProjectionUtils {
#define IS_EVAL3D 1
#endif
inline constexpr uint32_t NUM_WARPS = 10;
inline constexpr uint32_t NUM_THREADS = NUM_WARPS * WARP_SIZE;
// One CUDA block per micro-tile (TILE_SIZE_X x TILE_SIZE_Y pixels, one thread
// per pixel). Binning is still done at the coarser macro-tile granularity, so a
// block reads its macro tile's gaussian range and culls the ones that miss its
// micro-tile. Per-micro-tile blocks let the scheduler spread a dense macro
// tile's work across many SMs (a single huge macro tile no longer stalls the
// whole launch on one SM) and keep shared-memory tiny for high occupancy.
inline constexpr uint32_t TILE_AREA = TILE_SIZE_X * TILE_SIZE_Y;
template<
typename SplatPrimitive,
@@ -74,11 +79,14 @@ __global__ void rasterize_to_pixels_fwd_kernel(
RenderOutput::Buffer render_colors2, // [I, image_height, image_width, ...]
RenderOutput::Buffer render_distortions // [I, image_height, image_width, ...]
) {
uint32_t tid = threadIdx.x;
uint32_t wid = tid / WARP_SIZE;
uint32_t lid = tid % WARP_SIZE;
auto block = cg::this_thread_block();
uint32_t tid = threadIdx.x; // 0 .. TILE_AREA-1, one per pixel
int32_t image_id = blockIdx.x;
int32_t tile_id = blockIdx.y * tile_width + blockIdx.z;
// one block per micro-tile; recover the macro tile it belongs to for binning
uint32_t mt_y = blockIdx.y; // micro-tile row
uint32_t mt_x = blockIdx.z; // micro-tile col
int32_t tile_id = (mt_y / MACRO_TILE_SIZE_Y) * tile_width + (mt_x / MACRO_TILE_SIZE_X);
tile_offsets += image_id * tile_height * tile_width;
render_Ts += image_id * image_height * image_width;
@@ -100,146 +108,9 @@ __global__ void rasterize_to_pixels_fwd_kernel(
float3 ray_o = SlangProjectionUtils::transform_ray_o(R, t);
#endif
// number of warp-batches preloaded per depth batch, sized so the whole
// batch of gaussians fits in static shared memory (smaller for fatter
// fragments, e.g. 3DGUT, larger for lighter primitives)
constexpr uint32_t DEPTH_BATCH_SIZE = bool(IS_EVAL3D) ? 12 : 16;
constexpr uint32_t DEPTH_BATCH_SPLATS = DEPTH_BATCH_SIZE * WARP_SIZE;
// gaussians overlapping this macro tile, shared by all of its sub-tiles
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];
uint32_t num_batches =
(range_end - range_start + WARP_SIZE - 1) / WARP_SIZE;
uint32_t num_depth_batches = (num_batches + DEPTH_BATCH_SIZE - 1) / DEPTH_BATCH_SIZE;
num_depth_batches = max(num_depth_batches, 1);
// when there is more than one depth batch we round-trip per-pixel state
// (transmittance / colors / last index) through global memory between
// batches and have to remember which pixels already saturated.
const bool needs_resume = num_depth_batches > 1;
constexpr uint32_t MACRO_NUM_TILES = MACRO_TILE_SIZE_X * MACRO_TILE_SIZE_Y;
constexpr uint32_t WARPS_PER_TILE = (TILE_SIZE_X * TILE_SIZE_Y) / WARP_SIZE;
// a sub-tile is fully done once every one of its pixels is outside the image
// or has saturated; such sub-tiles are dropped from later depth batches and
// the remaining ones are compacted so warps never pick an idle sub-tile.
__shared__ bool tile_done[MACRO_NUM_TILES];
__shared__ int active_tiles[MACRO_NUM_TILES];
__shared__ int numActiveTiles;
__shared__ uint32_t numTilesProcessed;
// a whole depth batch of gaussians is loaded once and shared by every
// sub-tile of the macro tile (they all sweep the same gaussian range)
__shared__ typename SplatPrimitive::FragmentFwd splat_batch[DEPTH_BATCH_SPLATS];
// per-gaussian sub-tile intersection mask: bit s set => gaussian overlaps
// sub-tile s. MACRO_NUM_TILES == WARP_SIZE, so one uint32 holds all sub-tiles.
static_assert(MACRO_NUM_TILES <= 32);
__shared__ uint32_t isect_mask[DEPTH_BATCH_SPLATS];
if (tid < MACRO_NUM_TILES)
tile_done[tid] = false;
for (uint32_t depth_batch = 0; depth_batch < num_depth_batches; ++depth_batch) {
__syncthreads();
// compact the still-active sub-tiles into a dense list
if (tid == 0) {
int k = 0;
for (int i = 0; i < (int)MACRO_NUM_TILES; ++i)
if (!tile_done[i])
active_tiles[k++] = i;
numActiveTiles = k;
numTilesProcessed = 0;
}
// preload this depth batch's gaussians into shared memory once
uint32_t db_base = (uint32_t)range_start + depth_batch * DEPTH_BATCH_SPLATS;
uint32_t db_count = min((uint32_t)range_end - db_base, DEPTH_BATCH_SPLATS);
uint32_t num_batches_inner = (db_count + WARP_SIZE - 1) / WARP_SIZE;
for (uint32_t k = tid; k < db_count; k += NUM_THREADS) {
int32_t g = flatten_ids[db_base + k]; // flatten index in [I * N] or [nnz]
splat_batch[k].load(splat_wbuffer, splat_sbuffer,
gaussian_ids ? gaussian_ids[g] : g % N, g);
}
#if !IS_EVAL3D
__syncthreads(); // non-eval3d derives the mask from the just-loaded fragments
#endif
// precompute, for each gaussian in the batch, the set of sub-tiles its
// footprint overlaps. One warp per gaussian, lane == sub-tile index, so a
// single __ballot packs all MACRO_NUM_TILES(<=32) sub-tiles into one word.
for (uint32_t g = wid; g < db_count; g += NUM_WARPS) {
bool hit = false;
if (lid < MACRO_NUM_TILES) {
uint32_t sx = lid % MACRO_TILE_SIZE_X;
uint32_t sy = lid / MACRO_TILE_SIZE_X;
float tx0 = (float)((blockIdx.z * MACRO_TILE_SIZE_X + sx) * TILE_SIZE_X);
float ty0 = (float)((blockIdx.y * MACRO_TILE_SIZE_Y + sy) * TILE_SIZE_Y);
#if IS_EVAL3D
// 3dgut keeps the projected 2D conic in the screen "scale" channel
// and the ellipse center is the AABB center; same alpha-threshold
// contour test as 2D splatting, matching the tile intersector.
int32_t sid = flatten_ids[db_base + g];
float opac = splat_sbuffer.opacities(sid);
if (opac > ALPHA_THRESHOLD) {
float3 conic = splat_sbuffer.scales(sid);
float4 bb = aabb[sid];
float ecx = 0.5f * (bb.x + bb.z);
float ecy = 0.5f * (bb.y + bb.w);
float kk = 0.5f / __logf(opac / ALPHA_THRESHOLD);
float3 inv_cov = { conic.x * kk, conic.y * kk, conic.z * kk };
hit = ellipse_box_overlap_test(
inv_cov,
tx0 - ecx, tx0 + TILE_SIZE_X - ecx,
ty0 - ecy, ty0 + TILE_SIZE_Y - ecy
);
}
#else
// 2D splatting: exact ellipse vs sub-tile box at the alpha threshold
typename SplatPrimitive::FragmentFwd sp = splat_batch[g];
if (sp.opac > ALPHA_THRESHOLD) {
float kk = 0.5f / __logf(sp.opac / ALPHA_THRESHOLD);
float3 inv_cov = { sp.conic.x * kk, sp.conic.y * kk, sp.conic.z * kk };
hit = ellipse_box_overlap_test(
inv_cov,
tx0 - sp.xy.x, tx0 + TILE_SIZE_X - sp.xy.x,
ty0 - sp.xy.y, ty0 + TILE_SIZE_Y - sp.xy.y
);
}
#endif
}
uint32_t m = __ballot_sync(~0u, hit);
if (lid == 0)
isect_mask[g] = m;
}
__syncthreads();
for (;;) {
// each warp processes one sub-tile, queued so no warp idles
uint32_t slot = 0;
if (lid == 0)
slot = atomicAdd(&numTilesProcessed, 1);
slot = __shfl_sync(~0u, slot, 0);
if (slot >= (uint32_t)numActiveTiles)
break;
uint32_t tileIdx = (uint32_t)active_tiles[slot];
uint32_t tileOffsetY = blockIdx.y * MACRO_TILE_SIZE_Y + (tileIdx / MACRO_TILE_SIZE_X);
uint32_t tileOffsetX = blockIdx.z * MACRO_TILE_SIZE_X + (tileIdx % MACRO_TILE_SIZE_X);
// whether every pixel of this sub-tile is done after this depth batch
bool sub_tile_done = true;
#pragma unroll
for (int tr_batch = 0; tr_batch < WARPS_PER_TILE; ++tr_batch) {
int local_pid = tr_batch * WARP_SIZE + lid;
uint32_t i = tileOffsetY * TILE_SIZE_Y + local_pid / TILE_SIZE_X;
uint32_t j = tileOffsetX * TILE_SIZE_X + local_pid % TILE_SIZE_X;
// this thread's pixel within the micro-tile
uint32_t i = mt_y * TILE_SIZE_Y + tid / TILE_SIZE_X;
uint32_t j = mt_x * TILE_SIZE_X + tid % TILE_SIZE_X;
float px = (float)j + 0.5f;
float py = (float)i + 0.5f;
@@ -256,14 +127,8 @@ __global__ void rasterize_to_pixels_fwd_kernel(
#endif
bool done = !inside;
// a pixel that saturated in an earlier depth batch must not resume
// accumulating in this one (the gaussians here are further back)
bool saturated = false;
// 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.
int32_t pix_id = i * image_width + j;
int32_t pix_id_global = image_id * image_height * image_width + pix_id;
float T = 1.0f;
@@ -273,42 +138,83 @@ __global__ void rasterize_to_pixels_fwd_kernel(
RenderOutput pix_out = RenderOutput::zero();
RenderOutput pix2_out = RenderOutput::zero();
RenderOutput distortion_out = RenderOutput::zero();
if (depth_batch > 0 && inside) {
// reload state persisted at the end of the previous depth batch
T = render_Ts[pix_id];
int32_t s = last_ids[pix_id];
saturated = (s < 0);
cur_idx = saturated ? (uint32_t)(-s - 1) : (uint32_t)s;
done = done || saturated;
pix_out = render_colors.load<SplatPrimitive::pixelType>(pix_id_global);
if constexpr (RenderOutput::has_depth(SplatPrimitive::pixelType))
pix_out.depth *= (1.0f - T);
if constexpr (output_distortion) {
pix2_out = render_colors2.load<SplatPrimitive::pixelType>(pix_id_global);
distortion_out = render_distortions.load<SplatPrimitive::pixelType>(pix_id_global);
}
}
for (uint32_t inner_batch = 0; inner_batch < num_batches_inner; ++inner_batch) {
// end early if every pixel in this warp is done
if (__popc(__ballot_sync(~0u, done)) >= WARP_SIZE)
// gaussians overlapping this micro-tile's macro 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];
uint32_t num_batches =
(range_end - range_start + TILE_AREA - 1) / TILE_AREA;
// micro-tile pixel box, for the footprint cull
const float mt_bx0 = (float)(mt_x * TILE_SIZE_X);
const float mt_by0 = (float)(mt_y * TILE_SIZE_Y);
__shared__ typename SplatPrimitive::FragmentFwd splat_batch[TILE_AREA];
// whether the loaded gaussian's footprint reaches this micro-tile (uniform
// across the whole block, so skipping it is divergence-free)
__shared__ bool splat_hit[TILE_AREA];
// collect and process batches of gaussians, one thread loads one gaussian
for (uint32_t b = 0; b < num_batches; ++b) {
// resync all threads before the next batch overwrites shared, and stop
// once every pixel in the block is done
if (__syncthreads_count(done) >= (int)TILE_AREA)
break;
// gaussians for this batch are already in shared memory; read directly
uint32_t local_base = inner_batch * WARP_SIZE;
uint32_t batch_start = db_base + local_base; // global flatten index base
uint32_t batch_size = min((uint32_t)WARP_SIZE, db_count - local_base);
// skip gaussians whose footprint never reaches this sub-tile; the
// survivors are compacted via ballot so we only evaluate the hits
bool lane_hit = (lid < batch_size) &&
((isect_mask[local_base + lid] >> tileIdx) & 1u);
uint32_t surv = __ballot_sync(~0u, lane_hit);
while (surv) {
uint32_t t = __ffs(surv) - 1;
surv &= surv - 1;
uint32_t batch_start = range_start + TILE_AREA * b;
uint32_t idx = batch_start + tid;
bool hit = false;
if (idx < (uint32_t)range_end) {
int32_t g = flatten_ids[idx]; // flatten index in [I * N] or [nnz]
uint32_t wi = gaussian_ids ? gaussian_ids[g] : g % N;
#if IS_EVAL3D
// 3dgut keeps the projected 2D conic in the screen "scale" channel
// and the ellipse center is the AABB center. Cull here, before the
// (expensive) 3D fragment load, matching the tile intersector.
float opac = splat_sbuffer.opacities(g);
if (opac > ALPHA_THRESHOLD) {
float3 conic = splat_sbuffer.scales(g);
float4 bb = aabb[g];
float ecx = 0.5f * (bb.x + bb.z);
float ecy = 0.5f * (bb.y + bb.w);
float kk = 0.5f / __logf(opac / ALPHA_THRESHOLD);
float3 inv_cov = { conic.x * kk, conic.y * kk, conic.z * kk };
hit = ellipse_box_overlap_test(inv_cov,
mt_bx0 - ecx, mt_bx0 + TILE_SIZE_X - ecx,
mt_by0 - ecy, mt_by0 + TILE_SIZE_Y - ecy);
}
if (hit)
splat_batch[tid].load(splat_wbuffer, splat_sbuffer, wi, g);
#else
// 2D splatting: the fragment is just the screen data, so load it
// (cheap) and cull from its own conic / opacity.
splat_batch[tid].load(splat_wbuffer, splat_sbuffer, wi, g);
const typename SplatPrimitive::FragmentFwd& sp = splat_batch[tid];
if (sp.opac > ALPHA_THRESHOLD) {
float kk = 0.5f / __logf(sp.opac / ALPHA_THRESHOLD);
float3 inv_cov = { sp.conic.x * kk, sp.conic.y * kk, sp.conic.z * kk };
hit = ellipse_box_overlap_test(inv_cov,
mt_bx0 - sp.xy.x, mt_bx0 + TILE_SIZE_X - sp.xy.x,
mt_by0 - sp.xy.y, mt_by0 + TILE_SIZE_Y - sp.xy.y);
}
#endif
}
splat_hit[tid] = hit;
// wait for other threads to collect the gaussians in the batch
block.sync();
// process gaussians in the current batch for this pixel
uint32_t batch_size = min((uint32_t)TILE_AREA, (uint32_t)(range_end - batch_start));
for (uint32_t t = 0; t < batch_size; ++t) {
if (!splat_hit[t]) // gaussian misses this micro-tile (uniform skip)
continue;
if (done)
continue;
typename SplatPrimitive::FragmentFwd splat = splat_batch[local_base + t];
typename SplatPrimitive::FragmentFwd splat = splat_batch[t];
#if IS_EVAL3D
float alpha = splat.evaluate_alpha(ray_o, ray_d);
#else
@@ -339,24 +245,18 @@ __global__ void rasterize_to_pixels_fwd_kernel(
pix_out += color * vis;
cur_idx = batch_start + t;
T = next_T;
} else { done = true; saturated = true; }
} else done = true;
} // while (surv)
} // for (uint32_t t = 0; t < batch_size; ++t)
}
// this sub-tile is only finished once every lane is done
sub_tile_done &= (__popc(__ballot_sync(~0u, done)) >= WARP_SIZE);
if (i < image_height && j < image_width) {
render_Ts[pix_id] = T;
if constexpr (RenderOutput::has_depth(SplatPrimitive::pixelType))
pix_out.depth /= fmaxf(1.0f - T, 1e-10f);
pix_out.saveParamsToBuffer<SplatPrimitive::pixelType>(render_colors, pix_id_global);
// index in bin of last gaussian in this pixel; saturated pixels are
// flagged with a negative encoding so a later depth batch knows to stop
last_ids[pix_id] = (needs_resume && saturated)
? -static_cast<int32_t>(cur_idx) - 1
: static_cast<int32_t>(cur_idx);
// index in bin of last gaussian in this pixel
last_ids[pix_id] = static_cast<int32_t>(cur_idx);
// distortion
if constexpr (output_distortion) {
pix2_out.saveParamsToBuffer<SplatPrimitive::pixelType>(render_colors2, pix_id_global);
@@ -364,34 +264,6 @@ __global__ void rasterize_to_pixels_fwd_kernel(
}
}
} // for (int tr_batch = 0; tr_batch < WARPS_PER_TILE; ++tr_batch)
if (lid == 0)
tile_done[tileIdx] = sub_tile_done;
} // for (;;)
} // for (uint32_t depth_batch = 0; depth_batch < num_depth_batches; ++depth_batch)
// undo the saturation encoding so last_ids holds plain gaussian indices
if (needs_resume) {
__syncthreads();
constexpr uint32_t MACRO_PIX_H = MACRO_TILE_SIZE_Y * TILE_SIZE_Y;
constexpr uint32_t MACRO_PIX_W = MACRO_TILE_SIZE_X * TILE_SIZE_X;
uint32_t base_i = blockIdx.y * MACRO_PIX_H;
uint32_t base_j = blockIdx.z * MACRO_PIX_W;
for (uint32_t p = tid; p < MACRO_PIX_H * MACRO_PIX_W; p += NUM_THREADS) {
uint32_t i = base_i + p / MACRO_PIX_W;
uint32_t j = base_j + p % MACRO_PIX_W;
if (i < image_height && j < image_width) {
int32_t pix_id = i * image_width + j;
int32_t s = last_ids[pix_id];
if (s < 0)
last_ids[pix_id] = -s - 1;
}
}
}
}
template<
@@ -431,10 +303,10 @@ void rasterize_to_pixels_fwd_kernel_wrapper(
RenderOutput::Buffer render_colors2, // [I, image_height, image_width, ...]
RenderOutput::Buffer render_distortions // [I, image_height, image_width, ...]
) {
// Each block covers a tile on the image. In total there are
// I * tile_height * tile_width blocks.
dim3 threads = {NUM_THREADS, 1, 1};
dim3 grid = {I, tile_height, tile_width};
// One block per micro-tile. The macro tile spans MACRO_TILE_SIZE_{X,Y}
// micro-tiles, so the grid is the macro-tile grid scaled up accordingly.
dim3 threads = {TILE_AREA, 1, 1};
dim3 grid = {I, tile_height * MACRO_TILE_SIZE_Y, tile_width * MACRO_TILE_SIZE_X};
#if IS_EVAL3D
rasterize_to_pixels_eval3d_fwd_kernel<
+216
View File
@@ -118,6 +118,210 @@ def get_inputs(seed=44):
return inputs
# ---------------------------------------------------------------------------
# Representative synthetic scenes for profiling across gaussian distributions.
# These are parametric stand-ins that reproduce the *projected* spatial density
# and depth skew of real captures (not photometric realism). Each returns the
# same 8-tuple as get_inputs().
# ---------------------------------------------------------------------------
import math
def _look_at(eye, target, up=(0.0, 1.0, 0.0)):
eye = torch.tensor(eye, dtype=torch.float32)
target = torch.tensor(target, dtype=torch.float32)
up = torch.tensor(up, dtype=torch.float32)
fwd = target - eye
fwd = fwd / fwd.norm().clamp_min(1e-8) # camera looks along +z
right = torch.linalg.cross(up, fwd)
right = right / right.norm().clamp_min(1e-8)
tup = torch.linalg.cross(fwd, right)
R = torch.stack([right, tup, fwd], dim=0) # rows = camera axes in world
vm = torch.eye(4)
vm[:3, :3] = R
vm[:3, 3] = -(R @ eye)
return vm
def _plane(n, center, u, v, jitter, gen):
# n points on a planar slab: center + a*u + b*v + c*normal, a,b in [-1,1]
center = torch.tensor(center, dtype=torch.float32)
u = torch.tensor(u, dtype=torch.float32)
v = torch.tensor(v, dtype=torch.float32)
nrm = torch.linalg.cross(u, v)
nrm = nrm / nrm.norm().clamp_min(1e-8)
a = torch.rand(n, 1, generator=gen) * 2 - 1
b = torch.rand(n, 1, generator=gen) * 2 - 1
c = (torch.rand(n, 1, generator=gen) * 2 - 1) * jitter
return center + a * u + b * v + c * nrm
def _finish_scene(means, viewmats, scale_log_mean, opac_bias, seed):
n = means.shape[0]
g = torch.Generator().manual_seed(seed + 7)
quats = torch.nn.functional.normalize(torch.randn(n, 4, generator=g))
scales = torch.randn(n, 3, generator=g) * 0.5 + scale_log_mean
opacities = torch.randn(n, generator=g) + opac_bias
features_dc = torch.rand(n, 3, generator=g)
features_sh = 0.2 * torch.randn(n, (SH_DEGREE + 1) ** 2 - 1, 3, generator=g)
cx, cy = 0.5 * W, 0.5 * H
fx = fy = 0.4 * W
Ks = torch.tensor([[[fx, 0, cx], [0, fy, cy], [0, 0, 1]]]).repeat(viewmats.shape[0], 1, 1).float()
inputs = (means, quats, scales, opacities, features_dc, features_sh, viewmats, Ks)
inputs = [torch.nn.Parameter(x.contiguous().to(device).clone()) for x in inputs]
if WITH_UT:
inputs = inputs[:-2] + [viewmats.to(device), Ks.to(device)]
return inputs
def gen_scene(kind, n, seed=44):
g = torch.Generator().manual_seed(seed)
if kind == "uniform":
# baseline: gaussians spread through space, cameras around origin
means = torch.randn(n, 3, generator=g)
vms = [_look_at([3.0 * math.cos(2 * math.pi * k / B),
0.5 * (k - (B - 1) / 2),
3.0 * math.sin(2 * math.pi * k / B)], [0, 0, 0]) for k in range(B)]
return _finish_scene(means, torch.stack(vms), -4.5, 0.0, seed)
if kind == "object":
# small object-centered capture: dense compact blob + sparse far shell
n_obj = int(0.9 * n); n_bg = n - n_obj
obj = 0.5 * torch.randn(n_obj, 3, generator=g)
dirn = torch.nn.functional.normalize(torch.randn(n_bg, 3, generator=g))
bg = dirn * (8 + 4 * torch.rand(n_bg, 1, generator=g))
means = torch.cat([obj, bg], 0)
vms = [_look_at([3.2 * math.cos(2 * math.pi * k / B),
0.7 * (k - (B - 1) / 2),
3.2 * math.sin(2 * math.pi * k / B)], [0, 0, 0]) for k in range(B)]
return _finish_scene(means, torch.stack(vms), -5.3, 0.8, seed)
if kind == "street":
# narrow street: ground + two facades receding down +z, camera low.
# grazing ground + converging facades -> depth grows toward the horizon.
n3 = n // 3
ground = _plane(n3, [0, 0, 20], [6, 0, 0], [0, 0, 20], 0.05, g); ground[:, 1] = 0.0
facL = _plane(n3, [-3, 3, 20], [0, 3, 0], [0, 0, 20], 0.05, g)
facR = _plane(n - 2 * n3, [3, 3, 20], [0, 3, 0], [0, 0, 20], 0.05, g)
means = torch.cat([ground, facL, facR], 0)
vms = [_look_at([0.4 * (k - (B - 1) / 2), 1.2, -1.0 - 0.3 * k],
[0.4 * (k - (B - 1) / 2), 1.0, 20]) for k in range(B)]
return _finish_scene(means, torch.stack(vms), -4.8, 1.0, seed)
if kind == "indoor":
# multi-room: box of 6 planes, camera inside looking toward a corner so
# two walls recede at a grazing angle (deep) and two are frontal (shallow).
nf = n // 6
planes = [
_plane(nf, [0, -2, 0], [5, 0, 0], [0, 0, 5], 0.05, g), # floor
_plane(nf, [0, 3, 0], [5, 0, 0], [0, 0, 5], 0.05, g), # ceiling
_plane(nf, [-5, 0.5, 0], [0, 2.5, 0], [0, 0, 5], 0.05, g), # wall x-
_plane(nf, [5, 0.5, 0], [0, 2.5, 0], [0, 0, 5], 0.05, g), # wall x+
_plane(nf, [0, 0.5, -5], [5, 0, 0], [0, 2.5, 0], 0.05, g), # wall z-
_plane(n - 5 * nf, [0, 0.5, 5], [5, 0, 0], [0, 2.5, 0], 0.05, g), # wall z+
]
means = torch.cat(planes, 0)
vms = [_look_at([-2 + 0.5 * k, 0.0, -2 + 0.3 * k], [4, 0, 4]) for k in range(B)]
return _finish_scene(means, torch.stack(vms), -4.5, 0.5, seed)
if kind == "garden":
# unbounded outdoor: grazing ground + central foreground bushes + far shell
n_g = n // 3; n_f = n // 3; n_b = n - n_g - n_f
ground = _plane(n_g, [0, 0, 8], [12, 0, 0], [0, 0, 12], 0.03, g); ground[:, 1] = 0.0
fg = torch.stack([2.0 * torch.randn(n_f, generator=g),
0.6 * torch.rand(n_f, generator=g),
8 + 2.0 * torch.randn(n_f, generator=g)], dim=1)
dirn = torch.nn.functional.normalize(torch.randn(n_b, 3, generator=g))
shell = dirn.abs() * torch.tensor([30.0, 15.0, 30.0]) + torch.tensor([0.0, 2.0, 8.0])
means = torch.cat([ground, fg, shell], 0)
vms = [_look_at([0.6 * (k - (B - 1) / 2), 1.4, -1.0], [0, 0.6, 8]) for k in range(B)]
return _finish_scene(means, torch.stack(vms), -4.7, 0.7, seed)
raise ValueError(f"unknown scene {kind}")
@torch.no_grad()
def scene_tile_stats(means, viewmats, Ks):
# characterize skew: gaussian-center counts per tile (footprint ignored, so
# absolute counts under-read true binning, but the *skew* is representative)
m = means.detach().to(device).float()
out = {}
for name, (tx, ty) in [("macro", (64, 32)), ("micro", (8, 8))]:
nx = (W + tx - 1) // tx; ny = (H + ty - 1) // ty
mx = 0; nonempty = []
for c in range(viewmats.shape[0]):
R = viewmats[c, :3, :3].to(device).float(); t = viewmats[c, :3, 3].to(device).float()
cam = m @ R.T + t
z = cam[:, 2]
u = Ks[c, 0, 0].item() * cam[:, 0] / z + Ks[c, 0, 2].item()
v = Ks[c, 1, 1].item() * cam[:, 1] / z + Ks[c, 1, 2].item()
ok = (z > 1e-4) & (u >= 0) & (u < W) & (v >= 0) & (v < H)
idx = (v[ok] / ty).long() * nx + (u[ok] / tx).long()
cnt = torch.bincount(idx, minlength=nx * ny)
mx = max(mx, int(cnt.max().item()) if cnt.numel() else 0)
nz = cnt[cnt > 0]
if nz.numel():
nonempty.append(nz.float())
meanc = torch.cat(nonempty).mean().item() if nonempty else 0.0
out[name + "_max"] = mx
out[name + "_mean"] = meanc
return out
def _time_ms(fn, warmup=3, repeat=30):
from time import perf_counter
for _ in range(warmup):
fn()
torch.cuda.synchronize()
t0 = perf_counter()
for _ in range(repeat):
fn()
torch.cuda.synchronize()
return 1e3 * (perf_counter() - t0) / repeat
@pytest.mark.skipif(not torch.cuda.is_available(), reason="No CUDA device")
def profile_scenes(kinds=None, n=200000):
global N
N = n
if kinds is None:
kinds = ["uniform", "object", "garden", "street", "indoor"]
prim = "3dgut" if WITH_UT else ["3dgs", "mip"][IS_ANTIALIASED]
def _isect_vram_mb():
bd = _C.engine_get_pool_breakdown()
return sum(e[1] for e in bd if e[0].lower().startswith("isect.")) / 1e6
print(f"\n##### scene profiling (N={n}, primitive={prim}) #####")
print(f"{'scene':8s} | {'macro max/mean':>16s} | {'micro max/mean':>15s} | {'fwd ms':>7s} | {'bwd ms':>7s} | {'isectMB':>7s}")
for kind in kinds:
try:
inputs = gen_scene(kind, n)
st = scene_tile_stats(inputs[0], inputs[6], inputs[7])
cpu_inputs = [x.detach().cpu().contiguous() for x in inputs]
_, renderer = rasterize_ssplat(*cpu_inputs)
ft = _time_ms(lambda: _C.engine_forward_3dgs(
renderer.primitive, renderer.sh_degree_to_use, renderer.packed))
C, Hh, Ww = renderer.viewmats.shape[0], renderer.height, renderer.width
v_rgb = torch.randn(C, Hh, Ww, 3, device=device)
v_depth = torch.randn(C, Hh, Ww, 1, device=device)
v_Ts = torch.randn(C, Hh, Ww, 1, device=device)
bt = _time_ms(lambda: _C.engine_backward_from_render_grad(
renderer._tv(v_rgb), renderer._tv(v_depth), renderer._tv(v_Ts)))
vram = _isect_vram_mb()
print(f"{kind:8s} | {st['macro_max']:7d}/{st['macro_mean']:7.0f} | "
f"{st['micro_max']:6d}/{st['micro_mean']:5.1f} | {ft:7.2f} | {bt:7.2f} | {vram:7.1f}")
_C.engine_reset()
except Exception as e:
import traceback
traceback.print_exc()
try:
_C.engine_reset()
except Exception:
pass
@pytest.mark.skipif(not torch.cuda.is_available(), reason="No CUDA device")
def test_rasterization():
@@ -240,6 +444,18 @@ def profile_rasterization():
if __name__ == "__main__":
import itertools
import sys
# `python3 tests/test_rasterization.py scenes [ut]` -> benchmark the
# representative scenes instead of the correctness/profile suite.
if "scenes" in sys.argv:
PACKED = True; IS_FISHEYE = False
IS_ANTIALIASED = True; WITH_UT = False
profile_scenes()
if "ut" in sys.argv:
IS_ANTIALIASED = False; WITH_UT = True
profile_scenes()
sys.exit(0)
# packed x fisheye x (not_aa+no_ut, aa+no_ut, not_aa+with_ut)
modes = [(False, False), (True, False), (False, True)]