Skip to content

Commit 47d19fa

Browse files
committed
[Web] Let external owners defer device-loss disposal
Signed-off-by: Akaash Parthasarathy <akaashrp@gmail.com>
1 parent 016a5a0 commit 47d19fa

2 files changed

Lines changed: 53 additions & 3 deletions

File tree

‎web/src/runtime.ts‎

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -921,6 +921,7 @@ export class Instance implements Disposable {
921921
private initProgressCallback: Array<InitProgressCallback> = [];
922922
private rng: LinearCongruentialGenerator;
923923
private deviceLostIsError = true; // whether device.lost is due to actual error or dispose()
924+
private autoDisposeOnDeviceLost = true;
924925
private cacheState: CacheState = new CacheState();
925926

926927
/**
@@ -2080,7 +2081,7 @@ export class Instance implements Disposable {
20802081
});
20812082

20822083
device.lost.then((info: any) => {
2083-
if (this.deviceLostIsError) {
2084+
if (this.deviceLostIsError && this.autoDisposeOnDeviceLost) {
20842085
console.error("Device lost, calling Instance.dispose(). Please initialize again. ", info);
20852086
this.dispose();
20862087
}
@@ -2108,6 +2109,17 @@ export class Instance implements Disposable {
21082109
this.lib.webGPUContext = webGPUContext;
21092110
}
21102111

2112+
/**
2113+
* Configure automatic disposal after WebGPU device loss.
2114+
*
2115+
* External owners should disable automatic disposal after initialization if
2116+
* they serialize disposal with active runtime calls.
2117+
* @param enabled Whether device loss should immediately dispose this instance.
2118+
*/
2119+
setDeviceLostAutoDispose(enabled: boolean): void {
2120+
this.autoDisposeOnDeviceLost = enabled;
2121+
}
2122+
21112123
/** Register all object factory */
21122124
private registerObjectFactoryFuncs(): void {
21132125
this.registerObjectConstructor("ffi.Array",

‎web/tests/node/test_tensor_cache_webgpu.js‎

Lines changed: 40 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -39,7 +39,18 @@ function createInstance() {
3939
);
4040
}
4141

42-
function createMockGPUDevice({ detachWriteSources = false } = {}) {
42+
function createDeferred() {
43+
let resolve;
44+
const promise = new Promise((resolvePromise) => {
45+
resolve = resolvePromise;
46+
});
47+
return { promise, resolve };
48+
}
49+
50+
function createMockGPUDevice({
51+
detachWriteSources = false,
52+
lost = new Promise(() => {}),
53+
} = {}) {
4354
const buffers = [];
4455
const writes = [];
4556
const queue = {
@@ -65,7 +76,7 @@ function createMockGPUDevice({ detachWriteSources = false } = {}) {
6576
};
6677
const device = {
6778
queue,
68-
lost: new Promise(() => {}),
79+
lost,
6980
addEventListener: jest.fn(),
7081
pushErrorScope: jest.fn(),
7182
popErrorScope: jest.fn(() => Promise.resolve(null)),
@@ -94,6 +105,33 @@ function createArtifactCache(manifest, shard) {
94105
};
95106
}
96107

108+
test("an external owner can defer disposal after device loss", async () => {
109+
const lost = createDeferred();
110+
const tvm = createInstance();
111+
const dispose = jest.spyOn(tvm, "dispose");
112+
const gpu = createMockGPUDevice({ lost: lost.promise });
113+
const log = jest.spyOn(console, "error").mockImplementation(() => {});
114+
try {
115+
tvm.initWebGPU(gpu.device);
116+
tvm.setDeviceLostAutoDispose(false);
117+
118+
lost.resolve({ reason: "unknown", message: "test device loss" });
119+
await lost.promise;
120+
await Promise.resolve();
121+
122+
expect(dispose).not.toHaveBeenCalled();
123+
124+
tvm.dispose();
125+
expect(dispose).toHaveBeenCalledTimes(1);
126+
} finally {
127+
if (dispose.mock.calls.length === 0) {
128+
tvm.dispose();
129+
}
130+
dispose.mockRestore();
131+
log.mockRestore();
132+
}
133+
});
134+
97135
test("WebGPU tensor cache uploads pass-through records and decodes BF16 in place", async () => {
98136
const tvm = createInstance();
99137
const gpu = createMockGPUDevice();

0 commit comments

Comments
 (0)