UNPKG

playcanvas

Version:

Open-source WebGL/WebGPU 3D engine for the web

102 lines (86 loc) 4.16 kB
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; @group(0) @binding(0) var<storage, read> tileSplatCounts: array<u32>; @group(0) @binding(1) var<storage, read_write> smallTileList: array<u32>; @group(0) @binding(2) var<storage, read_write> largeTileList: array<u32>; @group(0) @binding(3) var<storage, read_write> rasterizeTileList: array<u32>; @group(0) @binding(4) var<storage, read_write> tileListCounts: array<atomic<u32>>; @group(0) @binding(5) var<storage, read_write> indirectDispatchArgs: array<u32>; @group(0) @binding(6) var<storage, read_write> largeTileOverflowBases: array<u32>; @group(0) @binding(8) var<storage, read_write> indirectDrawArgs: array<DrawIndirectArgs>; struct Uniforms { numTiles: u32, dispatchSlotOffset: u32, bufferCapacity: u32, maxWorkgroupsPerDim: u32, drawSlot: u32, } @group(0) @binding(7) var<uniform> uniforms: Uniforms; @compute @workgroup_size(256) fn main(@builtin(local_invocation_index) 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 };