playcanvas
Version:
Open-source WebGL/WebGPU 3D engine for the web
125 lines (123 loc) • 4.44 kB
JavaScript
const NUM_BUCKETS = 128;
const computeGsplatLocalBucketSortSource = `
const NUM_BUCKETS: u32 = ${NUM_BUCKETS}u;
const MAX_CHUNK_SIZE: u32 = 4096u;
const WG_SIZE: u32 = 256u;
var<storage, read_write> tileEntries: array<u32>;
var<storage, read> largeTileOverflowBases: array<u32>;
var<storage, read> tileSplatCounts: array<u32>;
var<storage, read> depthBuffer: array<u32>;
var<storage, read> largeTileList: array<u32>;
var<storage, read_write> chunkRanges: array<u32>;
var<storage, read_write> totalChunks: array<atomic<u32>>;
var<storage, read> tileListCounts: array<u32>;
struct Uniforms {
bufferCapacity: u32,
maxChunks: u32,
}
var<uniform> uniforms: Uniforms;
var<workgroup> sDepthMin: atomic<u32>;
var<workgroup> sDepthMax: atomic<u32>;
var<workgroup> sBucketCounts: array<atomic<u32>, NUM_BUCKETS>;
var<workgroup> sBucketOffsets: array<u32, NUM_BUCKETS + 1>;
var<workgroup> sBucketCursors: array<atomic<u32>, NUM_BUCKETS>;
fn main(
localIdx: u32,
wid: vec3u,
numWorkgroups: vec3u
) {
let largeTileIdx = wid.y * numWorkgroups.x + wid.x;
if (largeTileIdx >= tileListCounts[1]) {
return;
}
let tileIdx = largeTileList[largeTileIdx];
let tStart = tileSplatCounts[tileIdx];
let tEnd = tileSplatCounts[tileIdx + 1u];
let count = tEnd - tStart;
let overflowBase = largeTileOverflowBases[largeTileIdx];
if (overflowBase + count > uniforms.bufferCapacity) {
return;
}
if (localIdx == 0u) {
atomicStore(&sDepthMin, 0xFFFFFFFFu);
atomicStore(&sDepthMax, 0u);
}
if (localIdx < NUM_BUCKETS) {
atomicStore(&sBucketCounts[localIdx], 0u);
atomicStore(&sBucketCursors[localIdx], 0u);
}
workgroupBarrier();
for (var i: u32 = localIdx; i < count; i += WG_SIZE) {
let entryIdx = tileEntries[tStart + i];
let depthU = depthBuffer[entryIdx];
atomicMin(&sDepthMin, depthU);
atomicMax(&sDepthMax, depthU);
}
workgroupBarrier();
let depthMinU = atomicLoad(&sDepthMin);
let depthMaxU = atomicLoad(&sDepthMax);
let depthMin = bitcast<f32>(depthMinU);
let depthMax = bitcast<f32>(depthMaxU);
let logMin = log(max(depthMin, 1e-6));
let logRange = log(max(depthMax, 1e-6)) - logMin;
let bucketScale = select(f32(NUM_BUCKETS) / logRange, 0.0, logRange < 1e-10);
for (var i: u32 = localIdx; i < count; i += WG_SIZE) {
let entryIdx = tileEntries[tStart + i];
let depth = bitcast<f32>(depthBuffer[entryIdx]);
let bucket = min(u32((log(max(depth, 1e-6)) - logMin) * bucketScale), NUM_BUCKETS - 1u);
atomicAdd(&sBucketCounts[bucket], 1u);
tileEntries[overflowBase + i] = entryIdx;
}
workgroupBarrier();
if (localIdx == 0u) {
sBucketOffsets[0] = 0u;
for (var b: u32 = 0u; b < NUM_BUCKETS; b++) {
sBucketOffsets[b + 1u] = sBucketOffsets[b] + atomicLoad(&sBucketCounts[b]);
}
}
workgroupBarrier();
for (var i: u32 = localIdx; i < count; i += WG_SIZE) {
let entryIdx = tileEntries[overflowBase + i];
let depth = bitcast<f32>(depthBuffer[entryIdx]);
let bucket = min(u32((log(max(depth, 1e-6)) - logMin) * bucketScale), NUM_BUCKETS - 1u);
let writePos = sBucketOffsets[bucket] + atomicAdd(&sBucketCursors[bucket], 1u);
tileEntries[tStart + writePos] = entryIdx;
}
workgroupBarrier();
if (localIdx == 0u) {
var chunkStart: u32 = 0u;
var currentSize: u32 = 0u;
let maxChunks = uniforms.maxChunks;
for (var b: u32 = 0u; b < NUM_BUCKETS; b++) {
var bRemaining = sBucketOffsets[b + 1u] - sBucketOffsets[b];
if (bRemaining == 0u) {
continue;
}
while (bRemaining > 0u) {
let space = MAX_CHUNK_SIZE - currentSize;
let take = min(bRemaining, space);
currentSize += take;
bRemaining -= take;
if (currentSize == MAX_CHUNK_SIZE) {
let cIdx = atomicAdd(&totalChunks[0], 1u);
if (cIdx < maxChunks) {
chunkRanges[cIdx * 2u] = tStart + chunkStart;
chunkRanges[cIdx * 2u + 1u] = currentSize;
}
chunkStart += currentSize;
currentSize = 0u;
}
}
}
if (currentSize > 0u) {
let cIdx = atomicAdd(&totalChunks[0], 1u);
if (cIdx < maxChunks) {
chunkRanges[cIdx * 2u] = tStart + chunkStart;
chunkRanges[cIdx * 2u + 1u] = currentSize;
}
}
}
}
`;
export { computeGsplatLocalBucketSortSource };