playcanvas
Version:
Open-source WebGL/WebGPU 3D engine for the web
304 lines (303 loc) • 13.4 kB
JavaScript
import { Compute } from "../../platform/graphics/compute.js";
import { Shader } from "../../platform/graphics/shader.js";
import { StorageBuffer } from "../../platform/graphics/storage-buffer.js";
import { BindGroupFormat, BindStorageBufferFormat, BindUniformBufferFormat } from "../../platform/graphics/bind-group-format.js";
import { UniformBufferFormat, UniformFormat } from "../../platform/graphics/uniform-buffer-format.js";
import {
BUFFERUSAGE_COPY_DST,
BUFFERUSAGE_COPY_SRC,
SHADERLANGUAGE_WGSL,
SHADERSTAGE_COMPUTE,
UNIFORMTYPE_FLOAT,
UNIFORMTYPE_UINT,
UNIFORMTYPE_VEC3,
UNIFORMTYPE_VEC4,
UNIFORMTYPE_UVEC4
} from "../../platform/graphics/constants.js";
import { computeGsplatIntervalCullSource } from "../shader-lib/wgsl/chunks/gsplat/compute-gsplat-interval-cull.js";
import { computeGsplatIntervalScatterSource } from "../shader-lib/wgsl/chunks/gsplat/compute-gsplat-interval-scatter.js";
import { computeGsplatWriteIndirectArgsSource } from "../shader-lib/wgsl/chunks/gsplat/compute-gsplat-write-indirect-args.js";
import { PrefixSumKernel } from "../graphics/prefix-sum-kernel.js";
import { GSplatResourceBase } from "../gsplat/gsplat-resource-base.js";
const WORKGROUP_SIZE = 256;
const INDEX_COUNT = 6 * GSplatResourceBase.instanceSize;
const INTERVAL_STRIDE = 4;
class GSplatIntervalCompaction {
device;
compactedSplatIds = null;
intervalsBuffer = null;
countBuffer = null;
prefixSumKernel = null;
numSplatsBuffer = null;
sortElementCountBuffer = null;
allocatedCompactedCount = 0;
allocatedIntervalCount = 0;
allocatedCountBufferSize = 0;
_uploadedVersion = -1;
_cullComputePerspective = null;
_cullComputeFisheye = null;
_scatterCompute = null;
_writeIndirectArgsCompute = null;
_cullBindGroupFormatPerspective = null;
_cullBindGroupFormatFisheye = null;
_scatterBindGroupFormat = null;
_writeArgsBindGroupFormat = null;
_scatterUniformBufferFormat = null;
_writeArgsUniformBufferFormat = null;
constructor(device) {
this.device = device;
this.numSplatsBuffer = new StorageBuffer(device, 4, BUFFERUSAGE_COPY_SRC | BUFFERUSAGE_COPY_DST);
this.sortElementCountBuffer = new StorageBuffer(device, 4, BUFFERUSAGE_COPY_SRC | BUFFERUSAGE_COPY_DST);
this.prefixSumKernel = new PrefixSumKernel(device);
this._createUniformBufferFormats();
this._createScatterCompute();
this._createWriteIndirectArgsCompute();
}
destroy() {
this.compactedSplatIds?.destroy();
this.intervalsBuffer?.destroy();
this.countBuffer?.destroy();
this.prefixSumKernel?.destroy();
this.numSplatsBuffer?.destroy();
this.sortElementCountBuffer?.destroy();
this._destroyCullPass();
this._scatterCompute?.shader?.destroy();
this._scatterBindGroupFormat?.destroy();
this._writeIndirectArgsCompute?.shader?.destroy();
this._writeArgsBindGroupFormat?.destroy();
this.compactedSplatIds = null;
this.intervalsBuffer = null;
this.countBuffer = null;
this.prefixSumKernel = null;
this.numSplatsBuffer = null;
this.sortElementCountBuffer = null;
this._scatterCompute = null;
this._scatterBindGroupFormat = null;
this._writeIndirectArgsCompute = null;
this._writeArgsBindGroupFormat = null;
this._scatterUniformBufferFormat = null;
this._writeArgsUniformBufferFormat = null;
}
_destroyCullPass() {
this._cullComputePerspective?.shader?.destroy();
this._cullBindGroupFormatPerspective?.destroy();
this._cullComputePerspective = null;
this._cullBindGroupFormatPerspective = null;
this._cullComputeFisheye?.shader?.destroy();
this._cullBindGroupFormatFisheye?.destroy();
this._cullComputeFisheye = null;
this._cullBindGroupFormatFisheye = null;
}
_createUniformBufferFormats() {
const device = this.device;
this._scatterUniformBufferFormat = new UniformBufferFormat(device, [
new UniformFormat("numIntervals", UNIFORMTYPE_UINT),
new UniformFormat("pad0", UNIFORMTYPE_UINT),
new UniformFormat("pad1", UNIFORMTYPE_UINT),
new UniformFormat("pad2", UNIFORMTYPE_UINT)
]);
this._writeArgsUniformBufferFormat = new UniformBufferFormat(device, [
new UniformFormat("drawSlot", UNIFORMTYPE_UINT),
new UniformFormat("indexCount", UNIFORMTYPE_UINT),
new UniformFormat("dispatchSlotBase", UNIFORMTYPE_UINT),
new UniformFormat("totalSplats", UNIFORMTYPE_UINT),
new UniformFormat("sortIndirectInfo", UNIFORMTYPE_UVEC4)
]);
}
_createCullPass(fisheye) {
const device = this.device;
const suffix = fisheye ? "Fisheye" : "";
const bindGroupFormat = new BindGroupFormat(device, [
new BindUniformBufferFormat("uniforms", SHADERSTAGE_COMPUTE),
new BindStorageBufferFormat("intervals", SHADERSTAGE_COMPUTE, true),
new BindStorageBufferFormat("countBuffer", SHADERSTAGE_COMPUTE, false),
new BindStorageBufferFormat("boundsBuffer", SHADERSTAGE_COMPUTE, true),
new BindStorageBufferFormat("transformsBuffer", SHADERSTAGE_COMPUTE, true)
]);
const cdefines = /* @__PURE__ */ new Map([["{WORKGROUP_SIZE}", WORKGROUP_SIZE.toString()]]);
if (fisheye) {
cdefines.set("GSPLAT_FISHEYE", "");
}
const uniformBufferFormat = fisheye ? new UniformBufferFormat(device, [
new UniformFormat("cameraWorldPos", UNIFORMTYPE_VEC3),
new UniformFormat("maxTheta", UNIFORMTYPE_FLOAT),
new UniformFormat("cameraForward", UNIFORMTYPE_VEC3),
new UniformFormat("numIntervals", UNIFORMTYPE_UINT)
]) : new UniformBufferFormat(device, [
new UniformFormat("frustumPlanes", UNIFORMTYPE_VEC4, 6),
new UniformFormat("numIntervals", UNIFORMTYPE_UINT)
]);
const shader = new Shader(device, {
name: `GSplatIntervalCull${suffix}`,
shaderLanguage: SHADERLANGUAGE_WGSL,
cshader: computeGsplatIntervalCullSource,
cdefines,
computeBindGroupFormat: bindGroupFormat,
computeUniformBufferFormats: { uniforms: uniformBufferFormat }
});
const compute = new Compute(device, shader, `GSplatIntervalCull${suffix}`);
return { compute, bindGroupFormat };
}
_getCullCompute(fisheye) {
if (fisheye) {
if (!this._cullComputeFisheye) {
const { compute, bindGroupFormat } = this._createCullPass(true);
this._cullComputeFisheye = compute;
this._cullBindGroupFormatFisheye = bindGroupFormat;
}
return this._cullComputeFisheye;
}
if (!this._cullComputePerspective) {
const { compute, bindGroupFormat } = this._createCullPass(false);
this._cullComputePerspective = compute;
this._cullBindGroupFormatPerspective = bindGroupFormat;
}
return this._cullComputePerspective;
}
_createScatterCompute() {
const device = this.device;
this._scatterBindGroupFormat = new BindGroupFormat(device, [
new BindUniformBufferFormat("uniforms", SHADERSTAGE_COMPUTE),
new BindStorageBufferFormat("intervals", SHADERSTAGE_COMPUTE, true),
new BindStorageBufferFormat("prefixSumBuffer", SHADERSTAGE_COMPUTE, true),
new BindStorageBufferFormat("compactedOutput", SHADERSTAGE_COMPUTE, false)
]);
const cdefines = /* @__PURE__ */ new Map([
["{WORKGROUP_SIZE}", WORKGROUP_SIZE.toString()]
]);
const shader = new Shader(device, {
name: "GSplatIntervalScatter",
shaderLanguage: SHADERLANGUAGE_WGSL,
cshader: computeGsplatIntervalScatterSource,
cdefines,
computeBindGroupFormat: this._scatterBindGroupFormat,
computeUniformBufferFormats: { uniforms: this._scatterUniformBufferFormat }
});
this._scatterCompute = new Compute(device, shader, "GSplatIntervalScatter");
}
_createWriteIndirectArgsCompute() {
const device = this.device;
this._writeArgsBindGroupFormat = new BindGroupFormat(device, [
new BindStorageBufferFormat("prefixSumBuffer", SHADERSTAGE_COMPUTE, true),
new BindStorageBufferFormat("indirectDrawArgs", SHADERSTAGE_COMPUTE, false),
new BindStorageBufferFormat("numSplatsBuf", SHADERSTAGE_COMPUTE, false),
new BindStorageBufferFormat("indirectDispatchArgs", SHADERSTAGE_COMPUTE, false),
new BindStorageBufferFormat("sortElementCountBuf", SHADERSTAGE_COMPUTE, false),
new BindUniformBufferFormat("uniforms", SHADERSTAGE_COMPUTE)
]);
const cdefines = /* @__PURE__ */ new Map([
["{INSTANCE_SIZE}", GSplatResourceBase.instanceSize],
["{KEYGEN_THREADS_PER_WORKGROUP}", 256],
["{MAX_WORKGROUPS_PER_DIM}", device.limits.maxComputeWorkgroupsPerDimension || 65535]
]);
const shader = new Shader(device, {
name: "GSplatIntervalWriteIndirectArgs",
shaderLanguage: SHADERLANGUAGE_WGSL,
cshader: computeGsplatWriteIndirectArgsSource,
cdefines,
computeBindGroupFormat: this._writeArgsBindGroupFormat,
computeUniformBufferFormats: { uniforms: this._writeArgsUniformBufferFormat }
});
this._writeIndirectArgsCompute = new Compute(device, shader, "GSplatIntervalWriteIndirectArgs");
}
_ensureCapacity(numIntervals, totalActiveSplats) {
if (totalActiveSplats > this.allocatedCompactedCount) {
this.compactedSplatIds?.destroy();
this.allocatedCompactedCount = totalActiveSplats;
this.compactedSplatIds = new StorageBuffer(this.device, totalActiveSplats * 4, BUFFERUSAGE_COPY_SRC);
}
const requiredCountSize = numIntervals + 1;
if (requiredCountSize > this.allocatedCountBufferSize) {
this.countBuffer?.destroy();
this.allocatedCountBufferSize = requiredCountSize;
this.countBuffer = new StorageBuffer(this.device, requiredCountSize * 4);
if (this.prefixSumKernel) {
this.prefixSumKernel.destroyPasses();
}
}
}
uploadIntervals(worldState) {
if (worldState.version === this._uploadedVersion) return;
this._uploadedVersion = worldState.version;
const splats = worldState.splats;
const numIntervals = worldState.totalIntervals;
if (numIntervals === 0) return;
if (numIntervals > this.allocatedIntervalCount) {
this.intervalsBuffer?.destroy();
this.allocatedIntervalCount = numIntervals;
this.intervalsBuffer = new StorageBuffer(this.device, numIntervals * INTERVAL_STRIDE * 4, BUFFERUSAGE_COPY_DST);
}
const data = new Uint32Array(numIntervals * INTERVAL_STRIDE);
let writeIdx = 0;
for (let s = 0; s < splats.length; s++) {
const splat = splats[s];
if (splat.intervals.length > 0) {
const nodeIndices = splat.intervalNodeIndices;
for (let i = 0; i < splat.intervals.length; i += 2) {
const count = splat.intervals[i + 1] - splat.intervals[i];
data[writeIdx++] = splat.intervalOffsets[i / 2];
data[writeIdx++] = count;
data[writeIdx++] = splat.boundsBaseIndex + (nodeIndices.length > 0 ? nodeIndices[i / 2] : 0);
data[writeIdx++] = 0;
}
} else {
data[writeIdx++] = splat.intervalOffsets[0];
data[writeIdx++] = splat.activeSplats;
data[writeIdx++] = splat.boundsBaseIndex;
data[writeIdx++] = 0;
}
}
this.intervalsBuffer.write(0, data, 0, numIntervals * INTERVAL_STRIDE);
}
dispatchCompact(frustumCuller, numIntervals, totalActiveSplats, fisheyeEnabled) {
if (numIntervals === 0) return;
this._ensureCapacity(numIntervals, totalActiveSplats);
const cullCompute = this._getCullCompute(fisheyeEnabled);
cullCompute.setParameter("intervals", this.intervalsBuffer);
cullCompute.setParameter("countBuffer", this.countBuffer);
cullCompute.setParameter("boundsBuffer", frustumCuller.boundsBuffer);
cullCompute.setParameter("transformsBuffer", frustumCuller.transformsBuffer);
if (fisheyeEnabled) {
cullCompute.setParameter("cameraWorldPos", frustumCuller.fisheyeCameraPos);
cullCompute.setParameter("maxTheta", frustumCuller.fisheyeMaxTheta);
cullCompute.setParameter("cameraForward", frustumCuller.fisheyeCameraForward);
} else {
cullCompute.setParameter("frustumPlanes[0]", frustumCuller.frustumPlanes);
}
cullCompute.setParameter("numIntervals", numIntervals);
const cullWorkgroups = Math.ceil(numIntervals / WORKGROUP_SIZE);
cullCompute.setupDispatch(cullWorkgroups);
this.device.computeDispatch([cullCompute], "GSplatIntervalCull");
const prefixCount = numIntervals + 1;
this.prefixSumKernel.resize(this.countBuffer, prefixCount);
this.prefixSumKernel.dispatch(this.device);
const scatterCompute = this._scatterCompute;
scatterCompute.setParameter("intervals", this.intervalsBuffer);
scatterCompute.setParameter("prefixSumBuffer", this.countBuffer);
scatterCompute.setParameter("compactedOutput", this.compactedSplatIds);
scatterCompute.setParameter("numIntervals", numIntervals);
scatterCompute.setParameter("pad0", 0);
scatterCompute.setParameter("pad1", 0);
scatterCompute.setParameter("pad2", 0);
scatterCompute.setupDispatch(numIntervals);
this.device.computeDispatch([scatterCompute], "GSplatIntervalScatter");
}
writeIndirectArgs(drawSlot, dispatchSlotBase, numIntervals, sortIndirectInfo) {
const compute = this._writeIndirectArgsCompute;
compute.setParameter("prefixSumBuffer", this.countBuffer);
compute.setParameter("indirectDrawArgs", this.device.indirectDrawBuffer);
compute.setParameter("numSplatsBuf", this.numSplatsBuffer);
compute.setParameter("indirectDispatchArgs", this.device.indirectDispatchBuffer);
compute.setParameter("sortElementCountBuf", this.sortElementCountBuffer);
compute.setParameter("drawSlot", drawSlot);
compute.setParameter("indexCount", INDEX_COUNT);
compute.setParameter("dispatchSlotBase", dispatchSlotBase);
compute.setParameter("totalSplats", numIntervals);
compute.setParameter("sortIndirectInfo", sortIndirectInfo);
compute.setupDispatch(1);
this.device.computeDispatch([compute], "GSplatIntervalWriteIndirectArgs");
}
}
export {
GSplatIntervalCompaction
};