playcanvas
Version:
Open-source WebGL/WebGPU 3D engine for the web
191 lines (190 loc) • 6.35 kB
JavaScript
const computeGsplatLocalTileCountSource = `
#include "gsplatCommonCS"
#include "gsplatTileIntersectCS"
const MAX_TILE_ENTRIES: u32 = 0xFFFFu;
const BITMASK_W: u32 = 8u;
const BITMASK_H: u32 = 4u;
const LARGE_AABB_THRESHOLD: u32 = 64u;
var<storage, read> compactedSplatIds: array<u32>;
var<storage, read> sortElementCount: array<u32>;
var<storage, read_write> projCache: array<u32>;
var<storage, read_write> tileSplatCounts: array<atomic<u32>>;
struct Uniforms {
splatTextureSize: u32,
numTilesX: u32,
numTilesY: u32,
viewProj: mat4x4f,
viewMatrix: mat4x4f,
focal: f32,
viewportWidth: f32,
viewportHeight: f32,
nearClip: f32,
farClip: f32,
minPixelSize: f32,
isOrtho: u32,
exposure: f32,
alphaClip: f32,
minContribution: f32,
#ifdef GSPLAT_FISHEYE
fisheye_k: f32,
fisheye_inv_k: f32,
fisheye_projMat00: f32,
fisheye_projMat11: f32,
#endif
}
var<uniform> uniforms: Uniforms;
var<storage, read_write> pairBuffer: array<u32>;
var<storage, read_write> countersBuffer: array<atomic<u32>>;
var<storage, read_write> splatPairStart: array<u32>;
var<storage, read_write> splatPairCount: array<u32>;
var<storage, read_write> largeSplatIds: array<u32>;
var<storage, read_write> depthBuffer: array<u32>;
#include "gsplatComputeSplatCS"
#include "gsplatFormatDeclCS"
#include "gsplatFormatReadCS"
#include "gsplatProjectCommonCS"
fn main(
gid: vec3u,
numWorkgroups: vec3u
) {
let threadIdx = gid.y * (numWorkgroups.x * 256u) + gid.x;
let numVisible = sortElementCount[0];
let projected = projectSplatCommon(
threadIdx,
numVisible,
uniforms.alphaClip,
uniforms.minPixelSize,
uniforms.minContribution,
uniforms.viewMatrix,
uniforms.viewProj,
uniforms.focal,
uniforms.viewportWidth,
uniforms.viewportHeight,
uniforms.nearClip,
uniforms.farClip,
uniforms.isOrtho,
#ifdef GSPLAT_FISHEYE
uniforms.fisheye_k, uniforms.fisheye_inv_k,
uniforms.fisheye_projMat00, uniforms.fisheye_projMat11,
#endif
);
if (!projected.valid) {
if (threadIdx < numVisible) {
projCache[threadIdx * {CACHE_STRIDE}u + 6u] = 0u;
splatPairStart[threadIdx] = 0u;
splatPairCount[threadIdx] = 0u;
}
return;
}
let opacity = projected.opacity;
let proj = projected.proj;
let det = proj.a * proj.c - proj.b * proj.b;
let invDet = 1.0 / det;
let cx = 4.0 * proj.c * invDet;
let cy = -4.0 * proj.b * invDet;
let cz = 4.0 * proj.a * invDet;
let base = threadIdx * {CACHE_STRIDE}u;
projCache[base + 0u] = bitcast<u32>(proj.screen.x);
projCache[base + 1u] = bitcast<u32>(proj.screen.y);
projCache[base + 2u] = bitcast<u32>(cx);
projCache[base + 3u] = bitcast<u32>(cy);
projCache[base + 4u] = bitcast<u32>(cz);
#ifdef PICK_MODE
let pcIdVal = loadPcId().r;
projCache[base + 5u] = pcIdVal;
projCache[base + 6u] = pack2x16float(vec2f(0.0, opacity));
#else
let color = getColor();
var rgb = max(color, vec3f(0.0));
projCache[base + 5u] = pack2x16float(vec2f(rgb.x, rgb.y));
projCache[base + 6u] = pack2x16float(vec2f(rgb.z, opacity));
#endif
depthBuffer[threadIdx] = bitcast<u32>(proj.viewDepth);
let screen = proj.screen;
let eval = computeSplatTileEval(screen, cx, cy, cz, half(opacity),
uniforms.viewportWidth, uniforms.viewportHeight,
uniforms.alphaClip);
let radiusFactor = eval.radiusFactor;
projCache[base + 7u] = bitcast<u32>(-0.5 * radiusFactor);
let minTileX = max(0i, i32(floor(eval.splatMin.x / f32(TILE_SIZE))));
let maxTileX = min(i32(uniforms.numTilesX) - 1i, i32(floor(eval.splatMax.x / f32(TILE_SIZE))));
let minTileY = max(0i, i32(floor(eval.splatMin.y / f32(TILE_SIZE))));
let maxTileY = min(i32(uniforms.numTilesY) - 1i, i32(floor(eval.splatMax.y / f32(TILE_SIZE))));
let aabbW = u32(maxTileX - minTileX + 1i);
var deferredToLarge = false;
if (maxTileX >= minTileX && maxTileY >= minTileY &&
aabbW * u32(maxTileY - minTileY + 1i) > LARGE_AABB_THRESHOLD) {
let idx = atomicAdd(&countersBuffer[1], 1u);
if (idx < arrayLength(&largeSplatIds)) {
largeSplatIds[idx] = threadIdx;
deferredToLarge = true;
}
}
if (deferredToLarge) {
splatPairStart[threadIdx] = 0u;
splatPairCount[threadIdx] = 0u;
return;
}
var myPairCount: u32 = 0u;
var bitmask: u32 = 0u;
if (minTileX == maxTileX && minTileY == maxTileY) {
myPairCount = 1u;
bitmask = 1u;
} else {
for (var ty = minTileY; ty <= maxTileY; ty++) {
for (var tx = minTileX; tx <= maxTileX; tx++) {
let tMin = vec2f(f32(tx) * f32(TILE_SIZE), f32(ty) * f32(TILE_SIZE));
let tMax = tMin + vec2f(f32(TILE_SIZE));
if (tileIntersectsEllipse(tMin, tMax, screen, cx, cy, cz, radiusFactor)) {
myPairCount++;
let localX = u32(tx - minTileX);
let localY = u32(ty - minTileY);
if (localX < BITMASK_W && localY < BITMASK_H) {
let bitIdx = localY * BITMASK_W + localX;
bitmask |= (1u << bitIdx);
}
}
}
}
}
if (myPairCount == 0u) {
splatPairStart[threadIdx] = 0u;
splatPairCount[threadIdx] = 0u;
return;
}
let pairBase = atomicAdd(&countersBuffer[0], myPairCount);
splatPairStart[threadIdx] = pairBase;
splatPairCount[threadIdx] = myPairCount;
var j: u32 = 0u;
for (var ty = minTileY; ty <= maxTileY; ty++) {
for (var tx = minTileX; tx <= maxTileX; tx++) {
let localX = u32(tx - minTileX);
let localY = u32(ty - minTileY);
var hits: bool;
if (localX < BITMASK_W && localY < BITMASK_H) {
let bitIdx = localY * BITMASK_W + localX;
hits = (bitmask & (1u << bitIdx)) != 0u;
} else {
let tMin = vec2f(f32(tx) * f32(TILE_SIZE), f32(ty) * f32(TILE_SIZE));
let tMax = tMin + vec2f(f32(TILE_SIZE));
hits = tileIntersectsEllipse(tMin, tMax, screen, cx, cy, cz, radiusFactor);
}
if (hits) {
let tileIdx = u32(ty) * uniforms.numTilesX + u32(tx);
let localOff = atomicAdd(&tileSplatCounts[tileIdx], 1u);
if (localOff < MAX_TILE_ENTRIES) {
pairBuffer[pairBase + j] = (tileIdx << 16u) | (localOff & 0xFFFFu);
j++;
}
}
}
}
if (j != myPairCount) {
splatPairCount[threadIdx] = j;
}
}
`;
export {
computeGsplatLocalTileCountSource
};