Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion backend/accelerated/webgpu/web/package.json
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,8 @@
"build:wasm:plonk": "npm run build:wasm:assets && npm run build:wasm:plonk:webgpu && npm run build:wasm:plonk:native",
"build:wasm:plonk:native": "GOOS=js GOARCH=wasm go build -o dist/assets/plonk-native.wasm ../plonk/internal/wasmruntime/native",
"build:wasm:plonk:webgpu": "GOOS=js GOARCH=wasm go build -o dist/assets/plonk-webgpu.wasm ../plonk/internal/wasmruntime/webgpu",
"lint": "eslint ."
"lint": "eslint .",
"test": "npm run build && node --test test/*.test.mjs"
},
"devDependencies": {
"@eslint/js": "^9.0.0",
Expand Down
10 changes: 10 additions & 0 deletions backend/accelerated/webgpu/web/src/curvegpu/buffer_pool.ts
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@ export class BufferPool {
private readonly pool: Map<PoolKey, PoolEntry[]> = new Map();
private readonly meta = new WeakMap<GPUBuffer, { size: number; usage: number }>();
private totalBytes = 0;
private closed = false;

constructor(device: GPUDevice, options?: { maxPooledBytes?: number }) {
this.device = device;
Expand All @@ -39,6 +40,9 @@ export class BufferPool {
* May return a cached buffer from a previous `release` call.
*/
acquire(size: number, usage: number, label?: string): GPUBuffer {
if (this.closed) {
throw new Error("buffer pool is closed");
}
const roundedSize = nextPowerOfTwo(Math.max(4, size));
const key = poolKey(roundedSize, usage);
const entries = this.pool.get(key);
Expand All @@ -63,6 +67,11 @@ export class BufferPool {
buffer.destroy();
return;
}
if (this.closed) {
buffer.destroy();
this.meta.delete(buffer);
return;
}
if (this.totalBytes + m.size > this.maxBytes) {
buffer.destroy();
this.meta.delete(buffer);
Expand All @@ -83,6 +92,7 @@ export class BufferPool {
* is closed to avoid GPU memory leaks.
*/
destroy(): void {
this.closed = true;
for (const entries of this.pool.values()) {
for (const { buffer } of entries) {
buffer.destroy();
Expand Down
1 change: 1 addition & 0 deletions backend/accelerated/webgpu/web/src/curvegpu/context.ts
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,7 @@ export async function createCurveGPUContext(options: CurveGPUContextOptions = {}
}
closed = true;
bufferPool.destroy();
device.destroy();
},
};
}
25 changes: 17 additions & 8 deletions backend/accelerated/webgpu/web/src/curvegpu/msm_gpu_runtime.ts
Original file line number Diff line number Diff line change
Expand Up @@ -203,17 +203,26 @@ export async function readbackBuffer(
buffer: GPUBuffer,
size: number,
): Promise<Uint8Array> {
let mapped = false;
const staging = device.createBuffer({
label: "g1-readback-staging",
size,
usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ,
});
const encoder = device.createCommandEncoder({ label: "g1-readback-encoder" });
encoder.copyBufferToBuffer(buffer, 0, staging, 0, size);
device.queue.submit([encoder.finish()]);
await staging.mapAsync(GPUMapMode.READ);
const bytes = new Uint8Array(staging.getMappedRange()).slice();
staging.unmap();
staging.destroy();
return bytes;
try {
const encoder = device.createCommandEncoder({ label: "g1-readback-encoder" });
encoder.copyBufferToBuffer(buffer, 0, staging, 0, size);
device.queue.submit([encoder.finish()]);
await staging.mapAsync(GPUMapMode.READ);
mapped = true;
const bytes = new Uint8Array(staging.getMappedRange()).slice();
staging.unmap();
mapped = false;
return bytes;
} finally {
if (mapped) {
staging.unmap();
}
staging.destroy();
}
}
268 changes: 146 additions & 122 deletions backend/accelerated/webgpu/web/src/curvegpu/msm_pippenger.ts
Original file line number Diff line number Diff line change
Expand Up @@ -68,24 +68,46 @@ export function buildJacPippengerRuntime(
count,
labelPrefix,
} = options;
const weightedBucketOutput = createEmptyPointStorageBuffer(device, `${labelPrefix}-weighted-out`, bucketCountOut, pointBytes);
const weightParams = createParamsBuffer(device, `${labelPrefix}-weight-params`, uniformBytes, { count: bucketCountOut });
const weightBindGroup = createBindGroupForBuffers(device, kernels.weightJac, `${labelPrefix}-weight-bg`,
bucketOutput, zeroInput, weightedBucketOutput, weightParams, bucketValuesInput);
await submitKernel(device, kernels.weightJac, weightBindGroup, bucketCountOut, `${labelPrefix}-weight`, workgroupSize, debug);
const cleanupBuffers: GPUBuffer[] = [];
let windowOutput: GPUBuffer | undefined;
let succeeded = false;
try {
const weightedBucketOutput = createEmptyPointStorageBuffer(device, `${labelPrefix}-weighted-out`, bucketCountOut, pointBytes);
cleanupBuffers.push(weightedBucketOutput);
const weightParams = createParamsBuffer(device, `${labelPrefix}-weight-params`, uniformBytes, { count: bucketCountOut });
cleanupBuffers.push(weightParams);
const weightBindGroup = createBindGroupForBuffers(device, kernels.weightJac, `${labelPrefix}-weight-bg`,
bucketOutput, zeroInput, weightedBucketOutput, weightParams, bucketValuesInput);
await submitKernel(device, kernels.weightJac, weightBindGroup, bucketCountOut, `${labelPrefix}-weight`, workgroupSize, debug);

const windowSize = Math.max(1, count * metadata.numWindows) * pointBytes;
const windowUsage = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST;
const windowOutput = pool
? pool.acquire(windowSize, windowUsage, `${labelPrefix}-window-out`)
: createEmptyPointStorageBuffer(device, `${labelPrefix}-window-out`, count * metadata.numWindows, pointBytes);
const windowSize = Math.max(1, count * metadata.numWindows) * pointBytes;
const windowUsage = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST;
windowOutput = pool
? pool.acquire(windowSize, windowUsage, `${labelPrefix}-window-out`)
: createEmptyPointStorageBuffer(device, `${labelPrefix}-window-out`, count * metadata.numWindows, pointBytes);

const windowParams = createParamsBuffer(device, `${labelPrefix}-window-params`, uniformBytes, { count: count * metadata.numWindows });
const windowBindGroup = createBindGroupForBuffers(device, kernels.subsumJac, `${labelPrefix}-window-bg`,
weightedBucketOutput, zeroInput, windowOutput, windowParams, bucketValuesInput, windowStartsInput, windowCountsInput);
await submitKernel(device, kernels.subsumJac, windowBindGroup, count * metadata.numWindows * workgroupSize,
`${labelPrefix}-window`, workgroupSize, debug);
return { windowOutput, cleanupBuffers: [weightedBucketOutput, weightParams, windowParams] };
const windowParams = createParamsBuffer(device, `${labelPrefix}-window-params`, uniformBytes, { count: count * metadata.numWindows });
cleanupBuffers.push(windowParams);
const windowBindGroup = createBindGroupForBuffers(device, kernels.subsumJac, `${labelPrefix}-window-bg`,
weightedBucketOutput, zeroInput, windowOutput, windowParams, bucketValuesInput, windowStartsInput, windowCountsInput);
await submitKernel(device, kernels.subsumJac, windowBindGroup, count * metadata.numWindows * workgroupSize,
`${labelPrefix}-window`, workgroupSize, debug);
succeeded = true;
return { windowOutput, cleanupBuffers };
} finally {
if (!succeeded) {
if (windowOutput) {
if (pool) {
pool.release(windowOutput);
} else {
windowOutput.destroy();
}
}
for (let i = cleanupBuffers.length - 1; i >= 0; i -= 1) {
cleanupBuffers[i].destroy();
}
}
}
},
};
}
Expand Down Expand Up @@ -147,118 +169,120 @@ export async function runSparseSignedPippengerMSM(options: {
}
const storageInUsage = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST;
const storagePointUsage = GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC | GPUBufferUsage.COPY_DST;
const pooledBuffers: GPUBuffer[] = [];
const ownedBuffers: GPUBuffer[] = [];
const trackPooled = (buffer: GPUBuffer): GPUBuffer => {
pooledBuffers.push(buffer);
return buffer;
};
const trackOwned = (buffer: GPUBuffer): GPUBuffer => {
ownedBuffers.push(buffer);
return buffer;
};

const zeroSize = Math.max(4, pointBytes);
const zeroInput = pool
? pool.acquire(zeroSize, storageInUsage, `${labelPrefix}-zero`)
: createStorageBufferFromBytes(device, `${labelPrefix}-zero`, zeroPointBytes, pointBytes);
if (pool) {
device.queue.writeBuffer(zeroInput, 0, zeroPointBytes.buffer, zeroPointBytes.byteOffset, zeroPointBytes.byteLength);
}

const basesSize = Math.max(1, termsPerInstance * count) * pointBytes;
const basesInput = pool
? pool.acquire(basesSize, storageInUsage, `${labelPrefix}-bases`)
: createStorageBufferFromBytes(device, `${labelPrefix}-bases`, basesBytes, basesSize);
if (pool) {
device.queue.writeBuffer(basesInput, 0, basesBytes.buffer, basesBytes.byteOffset, basesBytes.byteLength);
}
const baseIndicesInput = createU32StorageBuffer(device, `${labelPrefix}-base-indices`, metadata.baseIndices);
const bucketPointersInput = createU32StorageBuffer(device, `${labelPrefix}-bucket-pointers`, metadata.bucketPointers);
const bucketSizesInput = createU32StorageBuffer(device, `${labelPrefix}-bucket-sizes`, metadata.bucketSizes);
try {
const zeroSize = Math.max(4, pointBytes);
const zeroInput = pool
? trackPooled(pool.acquire(zeroSize, storageInUsage, `${labelPrefix}-zero`))
: trackOwned(createStorageBufferFromBytes(device, `${labelPrefix}-zero`, zeroPointBytes, pointBytes));
if (pool) {
device.queue.writeBuffer(zeroInput, 0, zeroPointBytes.buffer, zeroPointBytes.byteOffset, zeroPointBytes.byteLength);
}

const bucketCountOut = metadata.bucketPointers.length;
const bucketSize = Math.max(1, bucketCountOut) * pointBytes;
const bucketOutput = pool
? pool.acquire(bucketSize, storagePointUsage, `${labelPrefix}-bucket-out`)
: createEmptyPointStorageBuffer(device, `${labelPrefix}-bucket-out`, bucketCountOut, pointBytes);
const bucketParams = createParamsBuffer(device, `${labelPrefix}-bucket-params`, uniformBytes, {
count: bucketCountOut,
termsPerInstance,
window,
numWindows: metadata.numWindows,
bucketCount: metadata.bucketCount,
});
const bucketBindGroup = createBindGroupForBuffers(
device,
runtime.bucket,
`${labelPrefix}-bucket-bg`,
basesInput,
zeroInput,
bucketOutput,
bucketParams,
baseIndicesInput,
bucketPointersInput,
bucketSizesInput,
);
await submitKernel(device, runtime.bucket, bucketBindGroup, bucketCountOut, `${labelPrefix}-bucket`, runtime.bucketWorkgroupSize ?? 64, debug);
const basesSize = Math.max(1, termsPerInstance * count) * pointBytes;
const basesInput = pool
? trackPooled(pool.acquire(basesSize, storageInUsage, `${labelPrefix}-bases`))
: trackOwned(createStorageBufferFromBytes(device, `${labelPrefix}-bases`, basesBytes, basesSize));
if (pool) {
device.queue.writeBuffer(basesInput, 0, basesBytes.buffer, basesBytes.byteOffset, basesBytes.byteLength);
}
const baseIndicesInput = trackOwned(createU32StorageBuffer(device, `${labelPrefix}-base-indices`, metadata.baseIndices));
const bucketPointersInput = trackOwned(createU32StorageBuffer(device, `${labelPrefix}-bucket-pointers`, metadata.bucketPointers));
const bucketSizesInput = trackOwned(createU32StorageBuffer(device, `${labelPrefix}-bucket-sizes`, metadata.bucketSizes));

const bucketValuesInput = createU32StorageBuffer(device, `${labelPrefix}-bucket-values`, metadata.bucketValues);
const windowStartsInput = createU32StorageBuffer(device, `${labelPrefix}-window-starts`, metadata.windowStarts);
const windowCountsInput = createU32StorageBuffer(device, `${labelPrefix}-window-counts`, metadata.windowCounts);
const bucketCountOut = metadata.bucketPointers.length;
const bucketSize = Math.max(1, bucketCountOut) * pointBytes;
const bucketOutput = pool
? trackPooled(pool.acquire(bucketSize, storagePointUsage, `${labelPrefix}-bucket-out`))
: trackOwned(createEmptyPointStorageBuffer(device, `${labelPrefix}-bucket-out`, bucketCountOut, pointBytes));
const bucketParams = trackOwned(createParamsBuffer(device, `${labelPrefix}-bucket-params`, uniformBytes, {
count: bucketCountOut,
termsPerInstance,
window,
numWindows: metadata.numWindows,
bucketCount: metadata.bucketCount,
}));
const bucketBindGroup = createBindGroupForBuffers(
device,
runtime.bucket,
`${labelPrefix}-bucket-bg`,
basesInput,
zeroInput,
bucketOutput,
bucketParams,
baseIndicesInput,
bucketPointersInput,
bucketSizesInput,
);
await submitKernel(device, runtime.bucket, bucketBindGroup, bucketCountOut, `${labelPrefix}-bucket`, runtime.bucketWorkgroupSize ?? 64, debug);

const { windowOutput, cleanupBuffers: windowReductionCleanup } = await runtime.reduceWindows({
device,
pool,
pointBytes,
uniformBytes,
zeroInput,
bucketOutput,
bucketCountOut,
bucketValuesInput,
windowStartsInput,
windowCountsInput,
metadata,
count,
labelPrefix,
});
const bucketValuesInput = trackOwned(createU32StorageBuffer(device, `${labelPrefix}-bucket-values`, metadata.bucketValues));
const windowStartsInput = trackOwned(createU32StorageBuffer(device, `${labelPrefix}-window-starts`, metadata.windowStarts));
const windowCountsInput = trackOwned(createU32StorageBuffer(device, `${labelPrefix}-window-counts`, metadata.windowCounts));

const finalSize = Math.max(1, count) * pointBytes;
const finalOutput = pool
? pool.acquire(finalSize, storagePointUsage, `${labelPrefix}-final-out`)
: createEmptyPointStorageBuffer(device, `${labelPrefix}-final-out`, count, pointBytes);
const finalParams = createParamsBuffer(device, `${labelPrefix}-final-params`, uniformBytes, {
count,
termsPerInstance,
window,
numWindows: metadata.numWindows,
bucketCount: metadata.bucketCount,
});
const finalBindGroup = createBindGroupForBuffers(
device,
runtime.combine,
`${labelPrefix}-final-bg`,
windowOutput,
zeroInput,
finalOutput,
finalParams,
);
await submitKernel(device, runtime.combine, finalBindGroup, count, `${labelPrefix}-final`, 64, debug);
const { windowOutput, cleanupBuffers: windowReductionCleanup } = await runtime.reduceWindows({
device,
pool,
pointBytes,
uniformBytes,
zeroInput,
bucketOutput,
bucketCountOut,
bucketValuesInput,
windowStartsInput,
windowCountsInput,
metadata,
count,
labelPrefix,
});
if (pool) {
trackPooled(windowOutput);
} else {
trackOwned(windowOutput);
}
windowReductionCleanup.forEach(trackOwned);

const result = await readbackBuffer(device, finalOutput, Math.max(1, count) * pointBytes);
const finalSize = Math.max(1, count) * pointBytes;
const finalOutput = pool
? trackPooled(pool.acquire(finalSize, storagePointUsage, `${labelPrefix}-final-out`))
: trackOwned(createEmptyPointStorageBuffer(device, `${labelPrefix}-final-out`, count, pointBytes));
const finalParams = trackOwned(createParamsBuffer(device, `${labelPrefix}-final-params`, uniformBytes, {
count,
termsPerInstance,
window,
numWindows: metadata.numWindows,
bucketCount: metadata.bucketCount,
}));
const finalBindGroup = createBindGroupForBuffers(
device,
runtime.combine,
`${labelPrefix}-final-bg`,
windowOutput,
zeroInput,
finalOutput,
finalParams,
);
await submitKernel(device, runtime.combine, finalBindGroup, count, `${labelPrefix}-final`, 64, debug);

if (pool) {
pool.release(zeroInput);
pool.release(basesInput);
pool.release(bucketOutput);
pool.release(windowOutput);
pool.release(finalOutput);
} else {
zeroInput.destroy();
basesInput.destroy();
bucketOutput.destroy();
windowOutput.destroy();
finalOutput.destroy();
return await readbackBuffer(device, finalOutput, Math.max(1, count) * pointBytes);
} finally {
if (pool) {
for (let i = pooledBuffers.length - 1; i >= 0; i -= 1) {
pool.release(pooledBuffers[i]);
}
}
for (let i = ownedBuffers.length - 1; i >= 0; i -= 1) {
ownedBuffers[i].destroy();
}
}
baseIndicesInput.destroy();
bucketPointersInput.destroy();
bucketSizesInput.destroy();
bucketValuesInput.destroy();
windowStartsInput.destroy();
windowCountsInput.destroy();
windowReductionCleanup.forEach((buffer) => buffer.destroy());
bucketParams.destroy();
finalParams.destroy();

return result;
}
Loading