playcanvas
Version:
Open-source WebGL/WebGPU 3D engine for the web
2 lines (1 loc) • 6.22 kB
TypeScript
export const computeGsplatLocalTileCountLargeSource: "\n\n#include \"gsplatCommonCS\"\n#include \"gsplatTileIntersectCS\"\n\nconst WG_SIZE: u32 = 256u;\nconst MAX_TILE_ENTRIES: u32 = 0xFFFFu;\n\n@group(0) @binding(0) var<storage, read> projCache: array<u32>;\n@group(0) @binding(1) var<storage, read_write> tileSplatCounts: array<atomic<u32>>;\n@group(0) @binding(2) var<storage, read_write> pairBuffer: array<u32>;\n@group(0) @binding(3) var<storage, read_write> countersBuffer: array<atomic<u32>>;\n@group(0) @binding(4) var<storage, read_write> splatPairStart: array<u32>;\n@group(0) @binding(5) var<storage, read_write> splatPairCount: array<u32>;\n@group(0) @binding(6) var<storage, read> largeSplatIds: array<u32>;\n\nstruct Uniforms {\n numTilesX: u32,\n numTilesY: u32,\n viewportWidth: f32,\n viewportHeight: f32,\n alphaClip: f32,\n}\n@group(0) @binding(7) var<uniform> uniforms: Uniforms;\n\nvar<workgroup> wgPairCounts: array<u32, WG_SIZE>;\nvar<workgroup> wgPairOffsets: array<u32, WG_SIZE>;\nvar<workgroup> wgBase: u32;\n\n@compute @workgroup_size(256)\nfn main(\n @builtin(workgroup_id) wgId: vec3u,\n @builtin(num_workgroups) numWorkgroups: vec3u,\n @builtin(local_invocation_index) lid: u32\n) {\n let largeSplatIdx = wgId.y * numWorkgroups.x + wgId.x;\n let count = min(atomicLoad(&countersBuffer[1]), arrayLength(&largeSplatIds));\n\n // atomicLoad is non-uniform per WGSL rules, so early return would make\n // subsequent workgroupBarrier calls non-uniform. Use an active flag instead;\n // inactive workgroups still participate in barriers but skip all real work.\n let isActive = largeSplatIdx < count;\n\n var threadIdx = u32(0);\n var minTileX = 0i;\n var maxTileX = 0i;\n var minTileY = 0i;\n var maxTileY = 0i;\n var aabbW = u32(0);\n var totalTiles = u32(0);\n var screen = vec2f(0.0);\n var cx = 0.0f;\n var cy = 0.0f;\n var cz = 0.0f;\n var radiusFactor = 0.0f;\n\n if (isActive) {\n threadIdx = largeSplatIds[largeSplatIdx];\n\n let cacheBase = threadIdx * {CACHE_STRIDE}u;\n screen = vec2f(bitcast<f32>(projCache[cacheBase + 0u]), bitcast<f32>(projCache[cacheBase + 1u]));\n cx = bitcast<f32>(projCache[cacheBase + 2u]);\n cy = bitcast<f32>(projCache[cacheBase + 3u]);\n cz = bitcast<f32>(projCache[cacheBase + 4u]);\n let opacity = unpack2x16float(projCache[cacheBase + 6u]).y;\n\n let eval = computeSplatTileEval(screen, cx, cy, cz, half(opacity),\n uniforms.viewportWidth, uniforms.viewportHeight,\n uniforms.alphaClip);\n radiusFactor = eval.radiusFactor;\n\n minTileX = max(0i, i32(floor(eval.splatMin.x / f32(TILE_SIZE))));\n maxTileX = min(i32(uniforms.numTilesX) - 1i, i32(floor(eval.splatMax.x / f32(TILE_SIZE))));\n minTileY = max(0i, i32(floor(eval.splatMin.y / f32(TILE_SIZE))));\n maxTileY = min(i32(uniforms.numTilesY) - 1i, i32(floor(eval.splatMax.y / f32(TILE_SIZE))));\n\n // Guard against degenerate AABBs where maxTile < minTile. This can happen\n // when capScale-driven radius shrinkage makes the tile-eval AABB smaller than\n // the frustum-cull AABB. The u32 cast of the negative difference would wrap\n // to ~4 billion, causing the tile loops to iterate for millions of iterations\n // per thread and hang the GPU.\n if (maxTileX >= minTileX && maxTileY >= minTileY) {\n aabbW = u32(maxTileX - minTileX + 1i);\n totalTiles = aabbW * u32(maxTileY - minTileY + 1i);\n }\n }\n\n // --- Phase 1: each thread counts its intersecting tiles ---\n var myHitCount: u32 = 0u;\n for (var i = lid; i < totalTiles; i += WG_SIZE) {\n let localX = i % aabbW;\n let localY = i / aabbW;\n let tx = minTileX + i32(localX);\n let ty = minTileY + i32(localY);\n let tMin = vec2f(f32(tx) * f32(TILE_SIZE), f32(ty) * f32(TILE_SIZE));\n let tMax = tMin + vec2f(f32(TILE_SIZE));\n if (tileIntersectsEllipse(tMin, tMax, screen, cx, cy, cz, radiusFactor)) {\n myHitCount++;\n }\n }\n\n // --- Workgroup prefix sum + global pair allocation ---\n wgPairCounts[lid] = myHitCount;\n workgroupBarrier();\n\n if (lid == 0u && isActive) {\n var sum: u32 = 0u;\n for (var i: u32 = 0u; i < WG_SIZE; i++) {\n wgPairOffsets[i] = sum;\n sum += wgPairCounts[i];\n }\n if (sum > 0u) {\n wgBase = atomicAdd(&countersBuffer[0], sum);\n } else {\n wgBase = 0u;\n }\n splatPairStart[threadIdx] = wgBase;\n splatPairCount[threadIdx] = sum | 0x80000000u;\n }\n workgroupBarrier();\n\n let myBase = wgBase + wgPairOffsets[lid];\n\n // --- Phase 2: write pairs with atomicAdd on tileSplatCounts ---\n var j: u32 = 0u;\n for (var i = lid; i < totalTiles; i += WG_SIZE) {\n let localX = i % aabbW;\n let localY = i / aabbW;\n let tx = minTileX + i32(localX);\n let ty = minTileY + i32(localY);\n let tMin = vec2f(f32(tx) * f32(TILE_SIZE), f32(ty) * f32(TILE_SIZE));\n let tMax = tMin + vec2f(f32(TILE_SIZE));\n if (tileIntersectsEllipse(tMin, tMax, screen, cx, cy, cz, radiusFactor)) {\n let tileIdx = u32(ty) * uniforms.numTilesX + u32(tx);\n let localOff = atomicAdd(&tileSplatCounts[tileIdx], 1u);\n if (localOff < MAX_TILE_ENTRIES) {\n pairBuffer[myBase + j] = (tileIdx << 16u) | (localOff & 0xFFFFu);\n j++;\n }\n }\n }\n\n // If any pairs were dropped by the cap, correct the stored count via workgroup sum.\n wgPairCounts[lid] = j;\n workgroupBarrier();\n if (lid == 0u && isActive) {\n var actualTotal: u32 = 0u;\n for (var i: u32 = 0u; i < WG_SIZE; i++) {\n actualTotal += wgPairCounts[i];\n }\n let storedCount = splatPairCount[threadIdx] & 0x7FFFFFFFu;\n if (actualTotal != storedCount) {\n splatPairCount[threadIdx] = actualTotal | 0x80000000u;\n }\n }\n}\n";