This is an automated email from the ASF dual-hosted git repository.
tlopex pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new 5d4dc87213 [Web] Avoid copies when uploading WASM memory to WebGPU
(#20165)
5d4dc87213 is described below
commit 5d4dc8721307cf9677becacf0395c19be130e206
Author: Akaash Parthasarathy <[email protected]>
AuthorDate: Sat Aug 22 11:55:55 2026 -0700
[Web] Avoid copies when uploading WASM memory to WebGPU (#20165)
Avoid an intermediate `Uint8Array` copy when uploading aligned data from
WASM memory to WebGPU.
- Add a bounds-checked borrowed view of Wasm linear memory
- Pass that view directly to `GPUQueue.writeBuffer`
- Preserve padding for unaligned transfers
- Round GPU allocations and readbacks to WebGPU’s four-byte copy
alignment
- Validate GPU copy offsets and handle zero-length readbacks
`writeBuffer` snapshots its source before returning, so the borrowed
WASM view does not need to remain valid after the call.
---
web/src/memory.ts | 28 ++++++++++++++
web/src/webgpu.ts | 46 ++++++++++++++++------
web/tests/node/test_memory.js | 63 ++++++++++++++++++++++++++++++
web/tests/node/test_webgpu.js | 89 ++++++++++++++++++++++++++++++++++++++++---
4 files changed, 209 insertions(+), 17 deletions(-)
diff --git a/web/src/memory.ts b/web/src/memory.ts
index 00dc17d412..5da8b59640 100644
--- a/web/src/memory.ts
+++ b/web/src/memory.ts
@@ -137,6 +137,34 @@ export class Memory {
result.set(this.viewU8.subarray(ptr, ptr + numBytes));
return result;
}
+ /**
+ * Return a borrowed view of raw bytes in Wasm memory.
+ *
+ * The returned view aliases the current WebAssembly.Memory buffer and must
+ * not be retained across a call that can grow the memory.
+ *
+ * @param ptr The head address.
+ * @param numBytes The number of bytes.
+ */
+ viewRawBytes(ptr: Pointer, numBytes: number): Uint8Array {
+ if (this.buffer != this.memory.buffer) {
+ this.updateViews();
+ }
+ if (!Number.isSafeInteger(ptr) || ptr < 0) {
+ throw new Error(`Invalid Wasm memory pointer: ${ptr}`);
+ }
+ if (!Number.isSafeInteger(numBytes) || numBytes < 0) {
+ throw new Error(`Invalid Wasm memory byte length: ${numBytes}`);
+ }
+ const end = ptr + numBytes;
+ if (!Number.isSafeInteger(end) || end > this.viewU8.byteLength) {
+ throw new Error(
+ `Wasm memory range [${ptr}, ${end}) exceeds memory size ` +
+ `${this.viewU8.byteLength}`,
+ );
+ }
+ return this.viewU8.subarray(ptr, end);
+ }
/**
* Load null-terminated C-string from ptr.
* @param ptr The head address
diff --git a/web/src/webgpu.ts b/web/src/webgpu.ts
index 60b553bdac..038f9accca 100644
--- a/web/src/webgpu.ts
+++ b/web/src/webgpu.ts
@@ -30,6 +30,23 @@ export interface GPUDeviceDetectOutput {
device: GPUDevice;
}
+function roundUpToFourBytes(nbytes: number): number {
+ if (!Number.isSafeInteger(nbytes) || nbytes < 0) {
+ throw new Error(`Invalid WebGPU buffer size: ${nbytes}`);
+ }
+ const aligned = Math.ceil(nbytes / 4) * 4;
+ if (!Number.isSafeInteger(aligned)) {
+ throw new Error(`WebGPU buffer size is too large to align: ${nbytes}`);
+ }
+ return aligned;
+}
+
+function validateWebGPUCopyOffset(offset: number, name: string): void {
+ if (!Number.isSafeInteger(offset) || offset < 0 || offset % 4 != 0) {
+ throw new Error(`${name} must be a nonnegative multiple of four:
${offset}`);
+ }
+}
+
/**
* DetectGPU device in the environment.
*/
@@ -914,16 +931,14 @@ export class WebGPUContext {
// DeviceAPI
private deviceAllocDataSpace(nbytes: number): GPUPointer {
- // allocate 0 bytes buffer as 1 bytes buffer.
- if (nbytes == 0) {
- nbytes = 1;
- }
+ // WebGPU buffer copies and queue writes operate in four-byte units.
+ const allocationBytes = Math.max(4, roundUpToFourBytes(nbytes));
const buffer = tryCreateBuffer(this.device, {
- size: nbytes,
+ size: allocationBytes,
usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC |
GPUBufferUsage.COPY_DST,
});
- this.currAllocatedBytes += nbytes;
- this.allAllocatedBytes += nbytes;
+ this.currAllocatedBytes += buffer.size;
+ this.allAllocatedBytes += buffer.size;
if (this.currAllocatedBytes > this.peakAllocatedBytes) {
this.peakAllocatedBytes = this.currAllocatedBytes;
}
@@ -951,11 +966,12 @@ export class WebGPUContext {
toOffset: number,
nbytes: number
): void {
+ validateWebGPUCopyOffset(toOffset, "WebGPU destination offset");
// Flush batched compute passes before writing to a GPU buffer,
// otherwise the write may be reordered before pending dispatches
// that read from the same buffer.
this.flushCommands();
- let rawBytes = this.memory.loadRawBytes(from, nbytes);
+ let rawBytes = this.memory.viewRawBytes(from, nbytes);
if (rawBytes.length % 4 !== 0) {
// writeBuffer requires length to be multiples of 4, so we pad here
const toPad = 4 - rawBytes.length % 4;
@@ -1016,9 +1032,15 @@ export class WebGPUContext {
to: Pointer,
nbytes: number
): void {
+ validateWebGPUCopyOffset(fromOffset, "WebGPU source offset");
// Flush batched compute passes before the readback copy.
this.flushCommands();
- const gpuTemp = this.getOrCreateReadStagingBuffer(nbytes);
+ if (nbytes == 0) {
+ this.memory.storeRawBytes(to, new Uint8Array(0));
+ return;
+ }
+ const copyBytes = roundUpToFourBytes(nbytes);
+ const gpuTemp = this.getOrCreateReadStagingBuffer(copyBytes);
const copyEncoder = this.device.createCommandEncoder();
copyEncoder.copyBufferToBuffer(
@@ -1026,14 +1048,14 @@ export class WebGPUContext {
fromOffset,
gpuTemp,
0,
- nbytes
+ copyBytes
);
const copyCommands = copyEncoder.finish();
this.device.queue.submit([copyCommands]);
const readPromise = gpuTemp.mapAsync(GPUMapMode.READ).then(() => {
- const data = gpuTemp.getMappedRange(0, nbytes);
- this.memory.storeRawBytes(to, new Uint8Array(data));
+ const data = gpuTemp.getMappedRange(0, copyBytes);
+ this.memory.storeRawBytes(to, new Uint8Array(data).subarray(0, nbytes));
this.recycleReadStagingBuffer(gpuTemp);
});
// Chain with any existing pending read so sync() awaits all of them.
diff --git a/web/tests/node/test_memory.js b/web/tests/node/test_memory.js
index 39de27835c..2daa9ff9b4 100644
--- a/web/tests/node/test_memory.js
+++ b/web/tests/node/test_memory.js
@@ -46,6 +46,69 @@ test("loadRawBytes preserves the requested length at the end
of memory", () => {
expect(Array.from(result)).toEqual([5, 6, 0, 0]);
});
+test("viewRawBytes returns a borrowed Wasm memory view", () => {
+ const wasmMemory = new WebAssembly.Memory({ initial: 1 });
+ const memory = new Memory(wasmMemory);
+ const source = new Uint8Array(wasmMemory.buffer, 8, 4);
+ source.set([1, 2, 3, 4]);
+
+ const result = memory.viewRawBytes(8, 4);
+
+ expect(Array.from(result)).toEqual([1, 2, 3, 4]);
+ expect(result.buffer).toBe(wasmMemory.buffer);
+ source[0] = 10;
+ result[1] = 20;
+ expect(Array.from(result)).toEqual([10, 20, 3, 4]);
+ expect(Array.from(source)).toEqual([10, 20, 3, 4]);
+});
+
+test("viewRawBytes refreshes its backing view after memory growth", () => {
+ const wasmMemory = new WebAssembly.Memory({ initial: 1, maximum: 2 });
+ const memory = new Memory(wasmMemory);
+ const oldBuffer = wasmMemory.buffer;
+
+ wasmMemory.grow(1);
+ const source = new Uint8Array(wasmMemory.buffer, 65536, 4);
+ source.set([5, 6, 7, 8]);
+ const result = memory.viewRawBytes(65536, 4);
+
+ expect(result.buffer).toBe(wasmMemory.buffer);
+ expect(result.buffer).not.toBe(oldBuffer);
+ expect(Array.from(result)).toEqual([5, 6, 7, 8]);
+});
+
+test("viewRawBytes supports shared Wasm memory", () => {
+ const wasmMemory = new WebAssembly.Memory({
+ initial: 1,
+ maximum: 2,
+ shared: true,
+ });
+ const memory = new Memory(wasmMemory);
+ const source = new Uint8Array(wasmMemory.buffer, 16, 4);
+ source.set([1, 2, 3, 4]);
+
+ const result = memory.viewRawBytes(16, 4);
+
+ expect(result.buffer).toBe(wasmMemory.buffer);
+ expect(result.buffer).toBeInstanceOf(SharedArrayBuffer);
+ expect(Array.from(result)).toEqual([1, 2, 3, 4]);
+});
+
+test.each([
+ [-1, 1, "pointer"],
+ [0.5, 1, "pointer"],
+ [Number.MAX_SAFE_INTEGER + 1, 1, "pointer"],
+ [0, -1, "byte length"],
+ [0, 0.5, "byte length"],
+ [0, Number.MAX_SAFE_INTEGER + 1, "byte length"],
+ [65536, 1, "exceeds memory size"],
+ [65535, 2, "exceeds memory size"],
+ [Number.MAX_SAFE_INTEGER, 1, "exceeds memory size"],
+])("viewRawBytes rejects invalid range (%p, %p)", (ptr, nbytes, message) => {
+ const memory = new Memory(new WebAssembly.Memory({ initial: 1 }));
+ expect(() => memory.viewRawBytes(ptr, nbytes)).toThrow(message);
+});
+
test("CachedCallStack commits a view of its cached bytes", () => {
const memory = {
wasm32: true,
diff --git a/web/tests/node/test_webgpu.js b/web/tests/node/test_webgpu.js
index dcd0be2e09..30f011d7e1 100644
--- a/web/tests/node/test_webgpu.js
+++ b/web/tests/node/test_webgpu.js
@@ -111,6 +111,7 @@ function createContext(deviceOptions) {
const gpu = createMockDevice(deviceOptions);
const memory = {
loadRawBytes: jest.fn(),
+ viewRawBytes: jest.fn(),
storeRawBytes: jest.fn(),
};
const context = new WebGPUContext(memory, gpu.device);
@@ -179,12 +180,16 @@ test("a host write flushes pending GPU copies before
writeBuffer", () => {
test("an aligned CPU to GPU copy writes the requested bytes", () => {
const { context, device, queue, memory, destination } = createContext();
const copyToGPU = context.getDeviceAPI("deviceCopyToGPU");
- const rawBytes = new Uint8Array([1, 2, 3, 4, 5, 6, 7, 8]);
- memory.loadRawBytes.mockReturnValue(rawBytes);
+ const wasmMemory = new WebAssembly.Memory({ initial: 1 });
+ const rawBytes = new Uint8Array(wasmMemory.buffer, 128, 8);
+ rawBytes.set([1, 2, 3, 4, 5, 6, 7, 8]);
+ memory.viewRawBytes.mockReturnValue(rawBytes);
copyToGPU(128, destination, 12, rawBytes.length);
- expect(memory.loadRawBytes).toHaveBeenCalledWith(128, rawBytes.length);
+ expect(memory.viewRawBytes).toHaveBeenCalledWith(128, rawBytes.length);
+ expect(memory.loadRawBytes).not.toHaveBeenCalled();
+ expect(queue.writeBuffer.mock.calls[0][2].buffer).toBe(wasmMemory.buffer);
expect(queue.writeBuffer).toHaveBeenCalledWith(
device.createBuffer.mock.results[1].value,
12,
@@ -198,11 +203,12 @@ test("an unaligned CPU to GPU copy pads the write to four
bytes", () => {
const { context, device, queue, memory, destination } = createContext();
const copyToGPU = context.getDeviceAPI("deviceCopyToGPU");
const rawBytes = new Uint8Array([1, 2, 3]);
- memory.loadRawBytes.mockReturnValue(rawBytes);
+ memory.viewRawBytes.mockReturnValue(rawBytes);
copyToGPU(256, destination, 4, rawBytes.length);
- expect(memory.loadRawBytes).toHaveBeenCalledWith(256, rawBytes.length);
+ expect(memory.viewRawBytes).toHaveBeenCalledWith(256, rawBytes.length);
+ expect(memory.loadRawBytes).not.toHaveBeenCalled();
expect(queue.writeBuffer).toHaveBeenCalledTimes(1);
const [buffer, toOffset, data, dataOffset, nbytes] =
queue.writeBuffer.mock.calls[0];
@@ -213,6 +219,79 @@ test("an unaligned CPU to GPU copy pads the write to four
bytes", () => {
expect(nbytes).toBe(4);
});
+test("a non-four-byte GPU allocation is rounded up for padded writes", () => {
+ const { context, device } = createContext();
+ const allocate = context.getDeviceAPI("deviceAllocDataSpace");
+
+ allocate(3);
+
+ expect(device.createBuffer).toHaveBeenLastCalledWith({
+ size: 4,
+ usage: GPUBufferUsage.STORAGE |
+ GPUBufferUsage.COPY_SRC |
+ GPUBufferUsage.COPY_DST,
+ });
+ expect(context.currAllocatedBytes).toBe(64 + 64 + 4);
+});
+
+test.each([
+ [-1, "destination offset"],
+ [0.5, "destination offset"],
+ [2, "destination offset"],
+])("a CPU to GPU copy rejects invalid offset %p", (offset, message) => {
+ const { context, memory, destination } = createContext();
+ const copyToGPU = context.getDeviceAPI("deviceCopyToGPU");
+
+ expect(() => copyToGPU(128, destination, offset, 4)).toThrow(message);
+ expect(memory.viewRawBytes).not.toHaveBeenCalled();
+});
+
+test.each([
+ [-1, "source offset"],
+ [0.5, "source offset"],
+ [2, "source offset"],
+])("a GPU readback rejects invalid offset %p", (offset, message) => {
+ const { context, source } = createContext();
+ const copyFromGPU = context.getDeviceAPI("deviceCopyFromGPU");
+
+ expect(() => copyFromGPU(source, offset, 128, 4)).toThrow(message);
+});
+
+test("an unaligned GPU readback copies four bytes and stores the logical
bytes", async () => {
+ const {
+ context,
+ device,
+ memory,
+ source,
+ } = createContext();
+ const copyFromGPU = context.getDeviceAPI("deviceCopyFromGPU");
+ device.createBuffer.mockImplementationOnce((descriptor) => {
+ const mappedData = new Uint8Array([1, 2, 3, 99]).buffer;
+ return {
+ size: descriptor.size,
+ destroy: jest.fn(),
+ mapAsync: jest.fn(() => Promise.resolve()),
+ getMappedRange: jest.fn(() => mappedData),
+ unmap: jest.fn(),
+ };
+ });
+
+ copyFromGPU(source, 0, 128, 3);
+ await context.sync();
+
+ const copyEncoder = device.createCommandEncoder.mock.results[0].value;
+ expect(copyEncoder.copyBufferToBuffer).toHaveBeenCalledWith(
+ device.createBuffer.mock.results[0].value,
+ 0,
+ device.createBuffer.mock.results[2].value,
+ 0,
+ 4,
+ );
+ expect(memory.storeRawBytes).toHaveBeenCalledTimes(1);
+ expect(memory.storeRawBytes.mock.calls[0][0]).toBe(128);
+ expect(Array.from(memory.storeRawBytes.mock.calls[0][1])).toEqual([1, 2, 3]);
+});
+
test("a GPU readback flushes pending copies before its own submission", async
() => {
const {
context,