playcanvas
Version:
Open-source WebGL/WebGPU 3D engine for the web
102 lines (86 loc) • 4.16 kB
JavaScript
import indirectCoreCS from "../common/comp/indirect-core.js";
import dispatchCoreCS from "../common/comp/dispatch-core.js";
const computeGsplatLocalClassifySource = (
/* wgsl */
`
${indirectCoreCS}
${dispatchCoreCS}
const MAX_TILE_ENTRIES: u32 = 4096u;
const CLASSIFY_WORKGROUP: u32 = 256u;
var<storage, read> tileSplatCounts: array<u32>;
var<storage, read_write> smallTileList: array<u32>;
var<storage, read_write> largeTileList: array<u32>;
var<storage, read_write> rasterizeTileList: array<u32>;
var<storage, read_write> tileListCounts: array<atomic<u32>>;
var<storage, read_write> indirectDispatchArgs: array<u32>;
var<storage, read_write> largeTileOverflowBases: array<u32>;
var<storage, read_write> indirectDrawArgs: array<DrawIndirectArgs>;
struct Uniforms {
numTiles: u32,
dispatchSlotOffset: u32,
bufferCapacity: u32,
maxWorkgroupsPerDim: u32,
drawSlot: u32,
}
var<uniform> uniforms: Uniforms;
fn main( localIdx: u32) {
let numTiles = uniforms.numTiles;
// Total tile entries from prefix sum \u2014 overflow scratch region starts here
let totalEntries = tileSplatCounts[numTiles];
for (var i: u32 = localIdx; i < numTiles; i += CLASSIFY_WORKGROUP) {
let tStart = tileSplatCounts[i];
let tEnd = tileSplatCounts[i + 1u];
let count = tEnd - tStart;
if (count == 0u || tEnd > uniforms.bufferCapacity) {
continue;
}
let rIdx = atomicAdd(&tileListCounts[2], 1u);
rasterizeTileList[rIdx] = i;
if (count <= MAX_TILE_ENTRIES) {
let sIdx = atomicAdd(&tileListCounts[0], 1u);
smallTileList[sIdx] = i;
} else {
// Large tile: claim overflow scratch in the shared tileEntries buffer.
// tileListCounts[3] tracks total overflow entries claimed across all large tiles.
// Bucket sort checks bounds and skips tiles whose overflow exceeds capacity.
let overflowOffset = atomicAdd(&tileListCounts[3], count);
let lIdx = atomicAdd(&tileListCounts[1], 1u);
largeTileList[lIdx] = i;
largeTileOverflowBases[lIdx] = totalEntries + overflowOffset;
}
}
workgroupBarrier();
// Thread 0 writes indirect dispatch args for passes 4a (small sort), 4b (bucket), 5 (rasterize).
// Uses balanced 2D dispatch to stay within maxComputeWorkgroupsPerDimension with minimal waste:
// y = ceil(count / maxDim), x = ceil(count / y). Waste is at most y-1 workgroups (typically 0-1).
if (localIdx == 0u) {
let smallCount = atomicLoad(&tileListCounts[0]);
let largeCount = atomicLoad(&tileListCounts[1]);
let rasterizeCount = atomicLoad(&tileListCounts[2]);
let off = uniforms.dispatchSlotOffset;
let maxDim = uniforms.maxWorkgroupsPerDim;
// Slot 0: small tile sort \u2014 1 workgroup per tile
let smallDim = calcDispatch2D(smallCount, maxDim);
indirectDispatchArgs[off + 0u] = smallDim.x;
indirectDispatchArgs[off + 1u] = smallDim.y;
indirectDispatchArgs[off + 2u] = 1u;
// Slot 1: bucket pre-sort \u2014 1 workgroup per large tile
let largeDim = calcDispatch2D(largeCount, maxDim);
indirectDispatchArgs[off + 3u] = largeDim.x;
indirectDispatchArgs[off + 4u] = largeDim.y;
indirectDispatchArgs[off + 5u] = 1u;
// Slot 2: rasterize \u2014 1 workgroup per non-empty tile
let rasterDim = calcDispatch2D(rasterizeCount, maxDim);
indirectDispatchArgs[off + 6u] = rasterDim.x;
indirectDispatchArgs[off + 7u] = rasterDim.y;
indirectDispatchArgs[off + 8u] = 1u;
// Indirect draw args for tile-based composite: 6 vertices per tile quad
indirectDrawArgs[uniforms.drawSlot] = DrawIndirectArgs(rasterizeCount * 6u, 1u, 0u, 0u, 0u);
}
}
`
);
export {
computeGsplatLocalClassifySource
};