Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 4 additions & 3 deletions web/src/cache_state.ts
Original file line number Diff line number Diff line change
Expand Up @@ -124,8 +124,8 @@ export class LRUCache<K, V> {
* 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.
Expand All @@ -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<string, Disposable>;

Expand Down
5 changes: 5 additions & 0 deletions web/src/ctypes.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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);
*/
Expand Down
30 changes: 28 additions & 2 deletions web/src/runtime.ts
Original file line number Diff line number Diff line change
Expand Up @@ -921,6 +921,7 @@ export class Instance implements Disposable {
private initProgressCallback: Array<InitProgressCallback> = [];
private rng: LinearCongruentialGenerator;
private deviceLostIsError = true; // whether device.lost is due to actual error or dispose()
private autoDisposeOnDeviceLost = true;
private cacheState: CacheState = new CacheState();

/**
Expand Down Expand Up @@ -1868,13 +1869,27 @@ export class Instance implements Disposable {
*/
makeShapeTuple(shape: Array<number>): 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.
Expand Down Expand Up @@ -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();
}
Expand Down Expand Up @@ -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",
Expand Down
18 changes: 18 additions & 0 deletions web/tests/node/test_object.js
Original file line number Diff line number Diff line change
Expand Up @@ -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/);
});
42 changes: 40 additions & 2 deletions web/tests/node/test_tensor_cache_webgpu.js
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {
Expand All @@ -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)),
Expand Down Expand Up @@ -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();
Expand Down
Loading