export const computeGsplatLocalTileCountSource: "\n\n#include \"gsplatCommonCS\"\n#include \"gsplatTileIntersectCS\"\n\nconst CACHE_STRIDE: u32 = 7u;\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\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    if (threadIdx >= numVisible) {\n        return;\n    }\n\n    let splatId = compactedSplatIds[threadIdx];\n    setSplat(splatId);\n    let center = getCenter();\n    let opacity = getOpacity();\n\n    if (opacity < uniforms.alphaClip) {\n        projCache[threadIdx * CACHE_STRIDE + 6u] = 0u;\n        splatPairStart[threadIdx] = 0u;\n        splatPairCount[threadIdx] = 0u;\n        return;\n    }\n\n    let rotation = half4(getRotation());\n    let scale = half3(getScale());\n\n    let proj = computeSplatCov(\n        center, rotation, scale,\n        uniforms.viewMatrix, uniforms.viewProj,\n        uniforms.focal, uniforms.viewportWidth, uniforms.viewportHeight,\n        uniforms.nearClip, uniforms.farClip, opacity, uniforms.minPixelSize,\n        uniforms.isOrtho, uniforms.alphaClip, uniforms.minContribution,\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 (!proj.valid) {\n        projCache[threadIdx * CACHE_STRIDE + 6u] = 0u;\n        splatPairStart[threadIdx] = 0u;\n        splatPairCount[threadIdx] = 0u;\n        return;\n    }\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;\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    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";
