mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
optimize rasterization forward
This commit is contained in:
@@ -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<
|
||||
|
||||
@@ -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)]
|
||||
|
||||
Reference in New Issue
Block a user