UNPKG

playcanvas

Version:

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

2 lines (1 loc) 10.6 kB
export const computeGsplatLocalTileCountSource: "\n\n#include \"gsplatCommonCS\"\n#include \"gsplatTileIntersectCS\"\n\n// Caps the 16-bit localOffset field in packed pairs (tileIdx << 16 | localOffset).\nconst MAX_TILE_ENTRIES: u32 = 0xFFFFu;\n\n// 8x4 = 32 bits fits in a single u32 bitmask. The bit index localY * 8 + localX\n// compiles to a pure shift (localY << 3 | localX), avoiding any multiply.\nconst BITMASK_W: u32 = 8u;\nconst BITMASK_H: u32 = 4u;\n\n// Splats whose AABB exceeds this many tiles are deferred to a cooperative\n// large-splat pass where 256 threads handle them in parallel, eliminating\n// the wavefront divergence that otherwise causes a long GPU tail.\nconst LARGE_AABB_THRESHOLD: u32 = 64u;\n\n@group(0) @binding(0) var<storage, read> compactedSplatIds: array<u32>;\n@group(0) @binding(1) var<storage, read> sortElementCount: array<u32>;\n@group(0) @binding(2) var<storage, read_write> projCache: array<u32>;\n@group(0) @binding(3) var<storage, read_write> tileSplatCounts: array<atomic<u32>>;\n\nstruct Uniforms {\n splatTextureSize: u32,\n numTilesX: u32,\n numTilesY: u32,\n viewProj: mat4x4f,\n viewMatrix: mat4x4f,\n focal: f32,\n viewportWidth: f32,\n viewportHeight: f32,\n nearClip: f32,\n farClip: f32,\n minPixelSize: f32,\n isOrtho: u32,\n exposure: f32,\n alphaClip: f32,\n minContribution: f32,\n #ifdef GSPLAT_FISHEYE\n fisheye_k: f32,\n fisheye_inv_k: f32,\n fisheye_projMat00: f32,\n fisheye_projMat11: f32,\n #endif\n}\n@group(0) @binding(4) var<uniform> uniforms: Uniforms;\n\n// Pair buffer bindings for the scatter-free approach.\n// pairBuffer stores packed (tileIdx << 16 | localOffset) per splat-tile intersection.\n// splatPairStart/splatPairCount let the PlaceEntries pass locate each splat's pairs.\n// countersBuffer packs two atomic counters: [0] = global pair counter, [1] = large splat count.\n@group(0) @binding(5) var<storage, read_write> pairBuffer: array<u32>;\n@group(0) @binding(6) var<storage, read_write> countersBuffer: array<atomic<u32>>;\n@group(0) @binding(7) var<storage, read_write> splatPairStart: array<u32>;\n@group(0) @binding(8) var<storage, read_write> splatPairCount: array<u32>;\n@group(0) @binding(9) var<storage, read_write> largeSplatIds: array<u32>;\n@group(0) @binding(10) var<storage, read_write> depthBuffer: array<u32>;\n\n#include \"gsplatComputeSplatCS\"\n#include \"gsplatFormatDeclCS\"\n#include \"gsplatFormatReadCS\"\n#include \"gsplatProjectCommonCS\"\n\n// NOTE on tile entry cap: if a tile exceeds MAX_TILE_ENTRIES (65535), the atomicAdd\n// count overcounts. Impact: the prefix sum allocates extra tileEntries slots that go\n// unwritten (wasting capacity), and the rasterize pass processes stale/zero entries in\n// those slots (minor visual artifacts). In practice, minContribution and minPixelSize\n// culling remove small/distant splats before tile counting, limiting per-tile density\n// when zoomed out and making overflow unlikely.\n\n@compute @workgroup_size(256)\nfn main(\n @builtin(global_invocation_id) gid: vec3u,\n @builtin(num_workgroups) numWorkgroups: vec3u\n) {\n let threadIdx = gid.y * (numWorkgroups.x * 256u) + gid.x;\n let numVisible = sortElementCount[0];\n\n let projected = projectSplatCommon(\n threadIdx,\n numVisible,\n uniforms.alphaClip,\n uniforms.minPixelSize,\n uniforms.minContribution,\n uniforms.viewMatrix,\n uniforms.viewProj,\n uniforms.focal,\n uniforms.viewportWidth,\n uniforms.viewportHeight,\n uniforms.nearClip,\n uniforms.farClip,\n uniforms.isOrtho,\n #ifdef GSPLAT_FISHEYE\n uniforms.fisheye_k, uniforms.fisheye_inv_k,\n uniforms.fisheye_projMat00, uniforms.fisheye_projMat11,\n #endif\n );\n\n if (!projected.valid) {\n if (threadIdx < numVisible) {\n projCache[threadIdx * {CACHE_STRIDE}u + 6u] = 0u;\n splatPairStart[threadIdx] = 0u;\n splatPairCount[threadIdx] = 0u;\n }\n return;\n }\n\n let opacity = projected.opacity;\n let proj = projected.proj;\n\n let det = proj.a * proj.c - proj.b * proj.b;\n let invDet = 1.0 / det;\n let cx = 4.0 * proj.c * invDet;\n let cy = -4.0 * proj.b * invDet;\n let cz = 4.0 * proj.a * invDet;\n\n let base = threadIdx * {CACHE_STRIDE}u;\n projCache[base + 0u] = bitcast<u32>(proj.screen.x);\n projCache[base + 1u] = bitcast<u32>(proj.screen.y);\n projCache[base + 2u] = bitcast<u32>(cx);\n projCache[base + 3u] = bitcast<u32>(cy);\n projCache[base + 4u] = bitcast<u32>(cz);\n\n#ifdef PICK_MODE\n let pcIdVal = loadPcId().r;\n projCache[base + 5u] = pcIdVal;\n projCache[base + 6u] = pack2x16float(vec2f(0.0, opacity));\n#else\n let color = getColor();\n var rgb = max(color, vec3f(0.0));\n projCache[base + 5u] = pack2x16float(vec2f(rgb.x, rgb.y));\n projCache[base + 6u] = pack2x16float(vec2f(rgb.z, opacity));\n#endif\n\n depthBuffer[threadIdx] = bitcast<u32>(proj.viewDepth);\n\n let screen = proj.screen;\n let eval = computeSplatTileEval(screen, cx, cy, cz, half(opacity),\n uniforms.viewportWidth, uniforms.viewportHeight,\n uniforms.alphaClip);\n let radiusFactor = eval.radiusFactor;\n\n // Per-splat power cutoff for the rasterize pass: the Gaussian exponent below which\n // the splat's contribution at a pixel drops below alphaClip. Equal to -radiusFactor / 2\n // = -log(opacity / alphaClip), clamped with radiusFactor. For high-opacity splats this\n // is -4 (matching the global cutoff); for low-opacity splats it's tighter, letting the\n // rasterize kernel skip exp() and the blend chain entirely for non-contributing pixels.\n projCache[base + 7u] = bitcast<u32>(-0.5 * radiusFactor);\n\n let minTileX = max(0i, i32(floor(eval.splatMin.x / f32(TILE_SIZE))));\n let maxTileX = min(i32(uniforms.numTilesX) - 1i, i32(floor(eval.splatMax.x / f32(TILE_SIZE))));\n let minTileY = max(0i, i32(floor(eval.splatMin.y / f32(TILE_SIZE))));\n let maxTileY = min(i32(uniforms.numTilesY) - 1i, i32(floor(eval.splatMax.y / f32(TILE_SIZE))));\n\n let aabbW = u32(maxTileX - minTileX + 1i);\n\n // Defer large splats to the cooperative large-splat pass where\n // 256 threads process them in parallel, avoiding wavefront divergence.\n // If the buffer overflows, fall through to normal single-thread processing.\n // Guard: when capScale shrinks the tile-eval radius below the frustum-cull\n // radius, maxTile can drop below minTile. The u32 cast of that negative\n // difference wraps to ~4 billion, falsely triggering the threshold.\n // The original loop handles this harmlessly (minTile > maxTile \u2192 0 iters),\n // so we must not classify these degenerate AABBs as large.\n var deferredToLarge = false;\n if (maxTileX >= minTileX && maxTileY >= minTileY &&\n aabbW * u32(maxTileY - minTileY + 1i) > LARGE_AABB_THRESHOLD) {\n let idx = atomicAdd(&countersBuffer[1], 1u);\n if (idx < arrayLength(&largeSplatIds)) {\n largeSplatIds[idx] = threadIdx;\n deferredToLarge = true;\n }\n }\n\n if (deferredToLarge) {\n splatPairStart[threadIdx] = 0u;\n splatPairCount[threadIdx] = 0u;\n return;\n }\n\n // =========================================================================\n // Phase 1: Count tiles + build bitmask (pure ALU, no atomics)\n // =========================================================================\n var myPairCount: u32 = 0u;\n var bitmask: u32 = 0u;\n\n if (minTileX == maxTileX && minTileY == maxTileY) {\n myPairCount = 1u;\n bitmask = 1u;\n } else {\n for (var ty = minTileY; ty <= maxTileY; ty++) {\n for (var tx = minTileX; tx <= maxTileX; tx++) {\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 myPairCount++;\n let localX = u32(tx - minTileX);\n let localY = u32(ty - minTileY);\n if (localX < BITMASK_W && localY < BITMASK_H) {\n let bitIdx = localY * BITMASK_W + localX;\n bitmask |= (1u << bitIdx);\n }\n }\n }\n }\n }\n\n if (myPairCount == 0u) {\n splatPairStart[threadIdx] = 0u;\n splatPairCount[threadIdx] = 0u;\n return;\n }\n\n // =========================================================================\n // Per-thread pair buffer reservation (no barrier, no shared memory)\n // =========================================================================\n let pairBase = atomicAdd(&countersBuffer[0], myPairCount);\n\n // =========================================================================\n // Phase 2: Write pairs using bitmask (all data in registers from Phase 1)\n // =========================================================================\n splatPairStart[threadIdx] = pairBase;\n splatPairCount[threadIdx] = myPairCount;\n\n var j: u32 = 0u;\n for (var ty = minTileY; ty <= maxTileY; ty++) {\n for (var tx = minTileX; tx <= maxTileX; tx++) {\n\n let localX = u32(tx - minTileX);\n let localY = u32(ty - minTileY);\n\n var hits: bool;\n if (localX < BITMASK_W && localY < BITMASK_H) {\n let bitIdx = localY * BITMASK_W + localX;\n hits = (bitmask & (1u << bitIdx)) != 0u;\n } else {\n let tMin = vec2f(f32(tx) * f32(TILE_SIZE), f32(ty) * f32(TILE_SIZE));\n let tMax = tMin + vec2f(f32(TILE_SIZE));\n hits = tileIntersectsEllipse(tMin, tMax, screen, cx, cy, cz, radiusFactor);\n }\n\n if (hits) {\n let tileIdx = u32(ty) * uniforms.numTilesX + u32(tx);\n let localOff = atomicAdd(&tileSplatCounts[tileIdx], 1u);\n if (localOff < MAX_TILE_ENTRIES) {\n pairBuffer[pairBase + j] = (tileIdx << 16u) | (localOff & 0xFFFFu);\n j++;\n }\n }\n }\n }\n\n if (j != myPairCount) {\n splatPairCount[threadIdx] = j;\n }\n}\n";