[js/webgpu] fix external buffer registration (#22254)

### Description

Fixes the problem of running into failure when GPU inputs shuffled
between iterations.
This commit is contained in:
Yulong Wang 2024-09-28 10:36:40 -07:00 committed by GitHub
parent 52a8c1cae8
commit 1bda91fc57
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 14 additions and 20 deletions

View file

@ -785,15 +785,20 @@ export class WebGpuBackend {
this.sessionExternalDataMapping.set(sessionId, sessionInputOutputMapping);
}
// the buffer may be user created, or managed by GPU data manager.
// The GPU data manager will not manage these buffers. we register them as external buffers.
//
// The map `sessionInputOutputMapping` is used to store the data ID and buffer for each input/output. Once a
// specific input/output is registered, the data ID will not change.
const previousBuffer = sessionInputOutputMapping.get(index);
const id = this.gpuDataManager.registerExternalBuffer(buffer, size, previousBuffer?.[1]);
const id = this.gpuDataManager.registerExternalBuffer(buffer, size, previousBuffer);
sessionInputOutputMapping.set(index, [id, buffer]);
return id;
}
unregisterBuffers(sessionId: number): void {
const sessionInputOutputMapping = this.sessionExternalDataMapping.get(sessionId);
if (sessionInputOutputMapping) {
sessionInputOutputMapping.forEach((bufferInfo) => this.gpuDataManager.unregisterExternalBuffer(bufferInfo[1]));
sessionInputOutputMapping.forEach((bufferInfo) => this.gpuDataManager.unregisterExternalBuffer(bufferInfo[0]));
this.sessionExternalDataMapping.delete(sessionId);
}
}

View file

@ -52,12 +52,12 @@ export interface GpuDataManager {
* GPU data manager only manages a mapping between the buffer and the GPU data ID. It will not manage the lifecycle of
* the external buffer.
*/
registerExternalBuffer(buffer: GPUBuffer, originalSize: number, previousBuffer?: GPUBuffer): number;
registerExternalBuffer(buffer: GPUBuffer, originalSize: number, previous?: [GpuDataId, GPUBuffer]): number;
/**
* unregister an external buffer for IO Binding.
*/
unregisterExternalBuffer(buffer: GPUBuffer): void;
unregisterExternalBuffer(id: GpuDataId): void;
/**
* destroy all gpu buffers.
@ -196,9 +196,6 @@ class GpuDataManagerImpl implements GpuDataManager {
// The reusable uniform buffers
private freeUniformBuffers: Map<number, GPUBuffer[]>;
// The external buffers registered users for IO Binding.
private externalBuffers: Map<GPUBuffer, GpuDataId>;
// The pendingBuffers for capture graph.
// a SessionID -> GPUBuffer[] mapping.
private capturedPendingBuffers: Map<number, GPUBuffer[]>;
@ -209,7 +206,6 @@ class GpuDataManagerImpl implements GpuDataManager {
this.freeUniformBuffers = new Map();
this.buffersForUploadingPending = [];
this.buffersPending = [];
this.externalBuffers = new Map();
this.capturedPendingBuffers = new Map();
for (const [key] of bucketFreelist) {
@ -284,14 +280,11 @@ class GpuDataManagerImpl implements GpuDataManager {
);
}
registerExternalBuffer(buffer: GPUBuffer, originalSize: number, previousBuffer?: GPUBuffer): number {
registerExternalBuffer(buffer: GPUBuffer, originalSize: number, previous?: [GpuDataId, GPUBuffer]): number {
let id: number | undefined;
if (previousBuffer) {
id = this.externalBuffers.get(previousBuffer);
if (id === undefined) {
throw new Error('previous buffer is not registered');
}
if (buffer === previousBuffer) {
if (previous) {
id = previous[0];
if (buffer === previous[1]) {
LOG_DEBUG(
'verbose',
() =>
@ -304,13 +297,11 @@ class GpuDataManagerImpl implements GpuDataManager {
throw new Error(`Registering a different external buffer under graph capture mode is not supported yet.
Please use the previous external buffer!`);
}
this.externalBuffers.delete(previousBuffer);
} else {
id = createNewGpuDataId();
}
this.storageCache.set(id, { gpuData: { id, type: GpuDataType.default, buffer }, originalSize });
this.externalBuffers.set(buffer, id);
LOG_DEBUG(
'verbose',
() => `[WebGPU] GpuDataManager.registerExternalBuffer(size=${originalSize}) => id=${id}, registered.`,
@ -318,11 +309,9 @@ class GpuDataManagerImpl implements GpuDataManager {
return id;
}
unregisterExternalBuffer(buffer: GPUBuffer): void {
const id = this.externalBuffers.get(buffer);
unregisterExternalBuffer(id: GpuDataId): void {
if (id !== undefined) {
this.storageCache.delete(id);
this.externalBuffers.delete(buffer);
LOG_DEBUG('verbose', () => `[WebGPU] GpuDataManager.unregisterExternalBuffer() => id=${id}`);
}
}