diff --git a/web/src/cache_state.ts b/web/src/cache_state.ts index c571a071b838..ab7c5fcabf95 100644 --- a/web/src/cache_state.ts +++ b/web/src/cache_state.ts @@ -124,8 +124,8 @@ export class LRUCache { * the JS→WASM FFI boundary each time. During LLM decode, the same shapes * repeat every token (e.g. [1,32,128]), so caching avoids thousands of * redundant FFI round-trips. - * - Invalidation: Never. Shape tuples are immutable value objects that - * remain valid for the lifetime of the TVM instance. + * - Invalidation: Cache entries may be evicted, but returned shape tuples + * hold independent references and remain valid for their caller's scope. * * Future additions (follow-up PR): * - **uniformCache**: Caches GPU uniform buffers keyed by content hash. @@ -142,7 +142,8 @@ export class CacheState { * Key: comma-separated dimension string, e.g. "1,32,128" * Value: TVM ShapeTuple object (Disposable) * - * Invalidation rule: None required — shape tuples are immutable. + * Eviction releases only the cache's reference. Shape tuples returned to + * callers have independent references and normal scope-managed lifetimes. */ readonly shapeCache: LRUCache; diff --git a/web/src/ctypes.ts b/web/src/ctypes.ts index 1f91779692ef..0333f056ddd4 100644 --- a/web/src/ctypes.ts +++ b/web/src/ctypes.ts @@ -176,6 +176,11 @@ export type FTVMFFIWasmFunctionCreate = ( */ export type FTVMFFIWasmFunctionDeleter = (self: Pointer) => void; +/** + * int TVMFFIObjectIncRef(TVMFFIObjectHandle obj); + */ +export type FTVMFFIObjectIncRef = (obj: Pointer) => number; + /** * int TVMFFIObjectDecRef(TVMFFIObjectHandle obj); */ diff --git a/web/src/runtime.ts b/web/src/runtime.ts index abc49be59490..996d7adc33c5 100644 --- a/web/src/runtime.ts +++ b/web/src/runtime.ts @@ -921,6 +921,7 @@ export class Instance implements Disposable { private initProgressCallback: Array = []; private rng: LinearCongruentialGenerator; private deviceLostIsError = true; // whether device.lost is due to actual error or dispose() + private autoDisposeOnDeviceLost = true; private cacheState: CacheState = new CacheState(); /** @@ -1868,13 +1869,27 @@ export class Instance implements Disposable { */ makeShapeTuple(shape: Array): TVMObject { const key = CacheState.computeShapeKey(shape); - return this.cacheState.shapeCache.get(key, () => { + const cachedTuple = this.cacheState.shapeCache.get(key, () => { const shapeArray = shape.map((value) => new Scalar(value, "int")); const tuple = this.ctx.makeShapeTuple(...shapeArray); // Detach from scope so the cached object survives across scopes. this.detachFromCurrentScope(tuple); return tuple; }) as TVMObject; + + // The cache owns its wrapper and may release it on eviction. Give the + // caller an independent strong reference with the usual scope lifetime. + const handle = cachedTuple.getHandle(); + this.lib.checkCall( + (this.lib.exports.TVMFFIObjectIncRef as ctypes.FTVMFFIObjectIncRef)(handle) + ); + const callerTuple = new TVMObject(handle, this.lib, this.ctx); + try { + return this.attachToCurrentScope(callerTuple); + } catch (err) { + callerTuple.dispose(); + throw err; + } } /** * Get type index from type key. @@ -2066,7 +2081,7 @@ export class Instance implements Disposable { }); device.lost.then((info: any) => { - if (this.deviceLostIsError) { + if (this.deviceLostIsError && this.autoDisposeOnDeviceLost) { console.error("Device lost, calling Instance.dispose(). Please initialize again. ", info); this.dispose(); } @@ -2094,6 +2109,17 @@ export class Instance implements Disposable { this.lib.webGPUContext = webGPUContext; } + /** + * Configure automatic disposal after WebGPU device loss. + * + * External owners should disable automatic disposal after initialization if + * they serialize disposal with active runtime calls. + * @param enabled Whether device loss should immediately dispose this instance. + */ + setDeviceLostAutoDispose(enabled: boolean): void { + this.autoDisposeOnDeviceLost = enabled; + } + /** Register all object factory */ private registerObjectFactoryFuncs(): void { this.registerObjectConstructor("ffi.Array", diff --git a/web/tests/node/test_object.js b/web/tests/node/test_object.js index 0e9e43097a9f..f797be4466d6 100644 --- a/web/tests/node/test_object.js +++ b/web/tests/node/test_object.js @@ -43,3 +43,21 @@ test("object", () => { assert(t1.getHandle() == t.getHandle()); }); }); + +test("shape cache does not invalidate caller-owned tuples", () => { + tvm.beginScope(); + const disposedTuple = tvm.makeShapeTuple([987654321, -1]); + disposedTuple.dispose(); + const cachedTuple = tvm.makeShapeTuple([987654321, -1]); + assert.doesNotThrow(() => cachedTuple.typeKey()); + + const evictedTuple = tvm.makeShapeTuple([987654321, 0]); + for (let i = 1; i <= 256; ++i) { + tvm.makeShapeTuple([987654321, i]); + } + assert.doesNotThrow(() => evictedTuple.typeKey()); + + tvm.endScope(); + assert.throws(() => cachedTuple.getHandle(), /already been disposed/); + assert.throws(() => evictedTuple.getHandle(), /already been disposed/); +}); diff --git a/web/tests/node/test_tensor_cache_webgpu.js b/web/tests/node/test_tensor_cache_webgpu.js index 3383a2e41eeb..b09250eb7ccf 100644 --- a/web/tests/node/test_tensor_cache_webgpu.js +++ b/web/tests/node/test_tensor_cache_webgpu.js @@ -39,7 +39,18 @@ function createInstance() { ); } -function createMockGPUDevice({ detachWriteSources = false } = {}) { +function createDeferred() { + let resolve; + const promise = new Promise((resolvePromise) => { + resolve = resolvePromise; + }); + return { promise, resolve }; +} + +function createMockGPUDevice({ + detachWriteSources = false, + lost = new Promise(() => {}), +} = {}) { const buffers = []; const writes = []; const queue = { @@ -65,7 +76,7 @@ function createMockGPUDevice({ detachWriteSources = false } = {}) { }; const device = { queue, - lost: new Promise(() => {}), + lost, addEventListener: jest.fn(), pushErrorScope: jest.fn(), popErrorScope: jest.fn(() => Promise.resolve(null)), @@ -94,6 +105,33 @@ function createArtifactCache(manifest, shard) { }; } +test("an external owner can defer disposal after device loss", async () => { + const lost = createDeferred(); + const tvm = createInstance(); + const dispose = jest.spyOn(tvm, "dispose"); + const gpu = createMockGPUDevice({ lost: lost.promise }); + const log = jest.spyOn(console, "error").mockImplementation(() => {}); + try { + tvm.initWebGPU(gpu.device); + tvm.setDeviceLostAutoDispose(false); + + lost.resolve({ reason: "unknown", message: "test device loss" }); + await lost.promise; + await Promise.resolve(); + + expect(dispose).not.toHaveBeenCalled(); + + tvm.dispose(); + expect(dispose).toHaveBeenCalledTimes(1); + } finally { + if (dispose.mock.calls.length === 0) { + tvm.dispose(); + } + dispose.mockRestore(); + log.mockRestore(); + } +}); + test("WebGPU tensor cache uploads pass-through records and decodes BF16 in place", async () => { const tvm = createInstance(); const gpu = createMockGPUDevice();