Skip to content

Commit 5ac9a7f

Browse files
authored
[Web] Give callers their own reference to cached shape tuples (#20501)
`makeShapeTuple` returns the object the shape cache owns. The cache is a 256-entry LRU that disposes what it evicts, so once a prefill touches more than 256 distinct shapes a caller can hold a tuple that has already been freed. WebLLM hits this on a Radeon iGPU as `Object has already been disposed` followed by a device hang (mlc-ai/web-llm#844). #20130 kept evicted tuples alive until teardown, whereas here the tuple is owned by the caller. 1. `makeShapeTuple` takes a reference on the cached handle and returns a new object attached to the current scope, so eviction only drops the cache's reference. Like every other function that returns a TVM object, it now requires an open scope. 2. `setDeviceLostAutoDispose(false)` lets an owner that already disposes the instance itself keep it from disposing on the device-lost promise A cache hit costs 15 to 30 ns more per call. Tests cover eviction while a caller holds the tuple and deferred disposal after device loss. --------- Signed-off-by: Akaash Parthasarathy <akaashrp@gmail.com>
1 parent ca49d57 commit 5ac9a7f

5 files changed

Lines changed: 95 additions & 7 deletions

File tree

‎web/src/cache_state.ts‎

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -124,8 +124,8 @@ export class LRUCache<K, V> {
124124
* the JS→WASM FFI boundary each time. During LLM decode, the same shapes
125125
* repeat every token (e.g. [1,32,128]), so caching avoids thousands of
126126
* redundant FFI round-trips.
127-
* - Invalidation: Never. Shape tuples are immutable value objects that
128-
* remain valid for the lifetime of the TVM instance.
127+
* - Invalidation: Cache entries may be evicted, but returned shape tuples
128+
* hold independent references and remain valid for their caller's scope.
129129
*
130130
* Future additions (follow-up PR):
131131
* - **uniformCache**: Caches GPU uniform buffers keyed by content hash.
@@ -142,7 +142,8 @@ export class CacheState {
142142
* Key: comma-separated dimension string, e.g. "1,32,128"
143143
* Value: TVM ShapeTuple object (Disposable)
144144
*
145-
* Invalidation rule: None required — shape tuples are immutable.
145+
* Eviction releases only the cache's reference. Shape tuples returned to
146+
* callers have independent references and normal scope-managed lifetimes.
146147
*/
147148
readonly shapeCache: LRUCache<string, Disposable>;
148149

‎web/src/ctypes.ts‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -176,6 +176,11 @@ export type FTVMFFIWasmFunctionCreate = (
176176
*/
177177
export type FTVMFFIWasmFunctionDeleter = (self: Pointer) => void;
178178

179+
/**
180+
* int TVMFFIObjectIncRef(TVMFFIObjectHandle obj);
181+
*/
182+
export type FTVMFFIObjectIncRef = (obj: Pointer) => number;
183+
179184
/**
180185
* int TVMFFIObjectDecRef(TVMFFIObjectHandle obj);
181186
*/

‎web/src/runtime.ts‎

Lines changed: 28 additions & 2 deletions
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
/**
@@ -1868,13 +1869,27 @@ export class Instance implements Disposable {
18681869
*/
18691870
makeShapeTuple(shape: Array<number>): TVMObject {
18701871
const key = CacheState.computeShapeKey(shape);
1871-
return this.cacheState.shapeCache.get(key, () => {
1872+
const cachedTuple = this.cacheState.shapeCache.get(key, () => {
18721873
const shapeArray = shape.map((value) => new Scalar(value, "int"));
18731874
const tuple = this.ctx.makeShapeTuple(...shapeArray);
18741875
// Detach from scope so the cached object survives across scopes.
18751876
this.detachFromCurrentScope(tuple);
18761877
return tuple;
18771878
}) as TVMObject;
1879+
1880+
// The cache owns its wrapper and may release it on eviction. Give the
1881+
// caller an independent strong reference with the usual scope lifetime.
1882+
const handle = cachedTuple.getHandle();
1883+
this.lib.checkCall(
1884+
(this.lib.exports.TVMFFIObjectIncRef as ctypes.FTVMFFIObjectIncRef)(handle)
1885+
);
1886+
const callerTuple = new TVMObject(handle, this.lib, this.ctx);
1887+
try {
1888+
return this.attachToCurrentScope(callerTuple);
1889+
} catch (err) {
1890+
callerTuple.dispose();
1891+
throw err;
1892+
}
18781893
}
18791894
/**
18801895
* Get type index from type key.
@@ -2066,7 +2081,7 @@ export class Instance implements Disposable {
20662081
});
20672082

20682083
device.lost.then((info: any) => {
2069-
if (this.deviceLostIsError) {
2084+
if (this.deviceLostIsError && this.autoDisposeOnDeviceLost) {
20702085
console.error("Device lost, calling Instance.dispose(). Please initialize again. ", info);
20712086
this.dispose();
20722087
}
@@ -2094,6 +2109,17 @@ export class Instance implements Disposable {
20942109
this.lib.webGPUContext = webGPUContext;
20952110
}
20962111

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+
20972123
/** Register all object factory */
20982124
private registerObjectFactoryFuncs(): void {
20992125
this.registerObjectConstructor("ffi.Array",

‎web/tests/node/test_object.js‎

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,3 +43,21 @@ test("object", () => {
4343
assert(t1.getHandle() == t.getHandle());
4444
});
4545
});
46+
47+
test("shape cache does not invalidate caller-owned tuples", () => {
48+
tvm.beginScope();
49+
const disposedTuple = tvm.makeShapeTuple([987654321, -1]);
50+
disposedTuple.dispose();
51+
const cachedTuple = tvm.makeShapeTuple([987654321, -1]);
52+
assert.doesNotThrow(() => cachedTuple.typeKey());
53+
54+
const evictedTuple = tvm.makeShapeTuple([987654321, 0]);
55+
for (let i = 1; i <= 256; ++i) {
56+
tvm.makeShapeTuple([987654321, i]);
57+
}
58+
assert.doesNotThrow(() => evictedTuple.typeKey());
59+
60+
tvm.endScope();
61+
assert.throws(() => cachedTuple.getHandle(), /already been disposed/);
62+
assert.throws(() => evictedTuple.getHandle(), /already been disposed/);
63+
});

‎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)