From 545445b7137c11be65dfcdbcb37da9d2aafe512a Mon Sep 17 00:00:00 2001 From: Akaash Parthasarathy Date: Mon, 20 Jul 2026 05:24:38 -0400 Subject: [PATCH 1/2] [Web] Expose RNG state for deterministic restore --- web/src/index.ts | 1 + web/src/runtime.ts | 21 +++++++- web/src/support.ts | 33 ++++++++++-- web/tests/node/test_random_generator.js | 71 +++++++++++++++++++++++++ 4 files changed, 121 insertions(+), 5 deletions(-) diff --git a/web/src/index.ts b/web/src/index.ts index 4de7e0930354..7cd05e7621b0 100644 --- a/web/src/index.ts +++ b/web/src/index.ts @@ -40,6 +40,7 @@ export { export { Disposable, LibraryProvider } from "./types"; export { RPCServer } from "./rpc_server"; export { assert, wasmPath, LinearCongruentialGenerator } from "./support"; +export type { RNGState } from "./support"; export { detectGPUDevice, GPUDeviceDetectOutput } from "./webgpu"; export { LRUCache, CacheState } from "./cache_state"; export { createPolyfillWASI } from "./compact"; diff --git a/web/src/runtime.ts b/web/src/runtime.ts index 078a0c7df21f..9c83a4d9acb6 100644 --- a/web/src/runtime.ts +++ b/web/src/runtime.ts @@ -23,7 +23,12 @@ import { Pointer, PtrOffset, SizeOf, TypeIndex } from "./ctypes"; import { Disposable } from "./types"; import { Memory, CachedCallStack } from "./memory"; -import { assert, StringToUint8Array, LinearCongruentialGenerator } from "./support"; +import { + assert, + StringToUint8Array, + LinearCongruentialGenerator, + RNGState, +} from "./support"; import { Environment } from "./environment"; import { AsyncifyHandler } from "./asyncify"; import { FunctionInfo, WebGPUContext } from "./webgpu"; @@ -1542,6 +1547,20 @@ export class Instance implements Disposable { this.rng.setSeed(seed); } + /** + * Get the state of the internal LinearCongruentialGenerator. + */ + getRNGState(): RNGState { + return this.rng.getState(); + } + + /** + * Restore the state of the internal LinearCongruentialGenerator. + */ + setRNGState(state: RNGState): void { + this.rng.setState(state); + } + /** * Sample index via top-p sampling. * diff --git a/web/src/support.ts b/web/src/support.ts index be85e85b7bab..95ee697628ab 100644 --- a/web/src/support.ts +++ b/web/src/support.ts @@ -80,24 +80,26 @@ export function wasmPath(): string { * Linear congruential generator for random number generating that can be seeded. * * Follows the implementation of `include/tvm/support/random_engine.h`, which follows the - * sepcification in https://en.cppreference.com/w/cpp/numeric/random/linear_congruential_engine. + * specification in https://en.cppreference.com/w/cpp/numeric/random/linear_congruential_engine. * * Note `Number.MAX_SAFE_INTEGER = 2^53 - 1`, and our intermediates are strictly less than 2^48. */ +export type RNGState = number; + export class LinearCongruentialGenerator { readonly modulus: number; readonly multiplier: number; readonly increment: number; - // Always within the range (0, 2^32 - 1) non-inclusive; if 0, will forever generate 0. + // Always within the range (0, modulus) non-inclusive; if 0, will forever generate 0. private rand_state: number; /** * Set modulus, multiplier, and increment. Initialize `rand_state` according to `Date.now()`. */ constructor() { - this.modulus = 2147483647; // 2^32 - 1 - this.multiplier = 48271; // between 2^15 and 2^16 + this.modulus = 2147483647; // 2^31 - 1 + this.multiplier = 48271; // between 2^15 and 2^16 this.increment = 0; this.setSeed(Date.now()); } @@ -119,6 +121,29 @@ export class LinearCongruentialGenerator { this.checkRandState(); } + /** + * Get the current generator state for deterministic restoration. + */ + getState(): RNGState { + return this.rand_state; + } + + /** + * Restore a state returned by `getState()`. + */ + setState(state: RNGState): void { + if (!Number.isInteger(state)) { + throw new Error("RNG state should be an integer."); + } + if (state <= 0 || state >= this.modulus) { + throw new Error( + `RNG state should be an integer in (0, ${this.modulus}).`, + ); + } + this.rand_state = state; + this.checkRandState(); + } + /** * Generate the next integer in the range (0, this.modulus) non-inclusive, updating `rand_state`. * diff --git a/web/tests/node/test_random_generator.js b/web/tests/node/test_random_generator.js index aefbf0f56fec..2a7c82e5d195 100644 --- a/web/tests/node/test_random_generator.js +++ b/web/tests/node/test_random_generator.js @@ -18,6 +18,23 @@ */ const tvmjs = require("../../dist"); +function createRNGOnlyInstance() { + const tvm = Object.create(tvmjs.Instance.prototype); + tvm.rng = new tvmjs.LinearCongruentialGenerator(); + tvm.empty = function () { + return { + copyFrom(input) { + return { + toArray() { + return input; + }, + }; + }, + }; + }; + return tvm; +} + test("Test coverage of [0,100] inclusive", () => { const covered = Array(100); const rng = new tvmjs.LinearCongruentialGenerator(); @@ -44,6 +61,40 @@ test("Test whether the same seed make two RNGs generate same results", () => { } }); +test("Restoring RNG state reproduces next random floats", () => { + const rng1 = new tvmjs.LinearCongruentialGenerator(); + const rng2 = new tvmjs.LinearCongruentialGenerator(); + rng1.setSeed(42); + for (let i = 0; i < 8; i++) { + rng1.randomFloat(); + } + + const state = rng1.getState(); + const expected = []; + for (let i = 0; i < 16; i++) { + expected.push(rng1.randomFloat()); + } + + rng2.setState(state); + const restored = []; + for (let i = 0; i < expected.length; i++) { + restored.push(rng2.randomFloat()); + } + expect(restored).toEqual(expected); +}); + +test("Instance RNG state reproduces next uniform samples", () => { + const tvm = createRNGOnlyInstance(); + tvm.setSeed(123); + const state = tvm.getRNGState(); + const expected = Array.from(tvm.uniform([16], -1.0, 1.0, {}).toArray()); + + tvm.uniform([16], -1.0, 1.0, {}); + tvm.setRNGState(state); + const restored = Array.from(tvm.uniform([16], -1.0, 1.0, {}).toArray()); + expect(restored).toEqual(expected); +}); + test("Test two RNGs with different seeds generate different results", () => { const rng1 = new tvmjs.LinearCongruentialGenerator(); const rng2 = new tvmjs.LinearCongruentialGenerator(); @@ -67,3 +118,23 @@ test('Illegal argument to `setSeed()`', () => { rng1.setSeed(42.5); }).toThrow("Seed should be an integer."); }); + +test("Illegal argument to `setState()`", () => { + const rng = new tvmjs.LinearCongruentialGenerator(); + + expect(() => { + rng.setState(undefined); + }).toThrow("RNG state should be an integer."); + expect(() => { + rng.setState({}); + }).toThrow("RNG state should be an integer."); + expect(() => { + rng.setState(0); + }).toThrow("RNG state should be an integer in"); + expect(() => { + rng.setState(rng.modulus); + }).toThrow("RNG state should be an integer in"); + expect(() => { + rng.setState(1.5); + }).toThrow("RNG state should be an integer."); +}); From e58d1cbed1d3f25e2d39551c350a0abefd1abcdf Mon Sep 17 00:00:00 2001 From: Akaash Parthasarathy Date: Tue, 21 Jul 2026 04:01:27 -0400 Subject: [PATCH 2/2] [Web] Simplify RNG state restore checks --- web/src/support.ts | 1 - web/tests/node/test_random_generator.js | 29 ------------------------- 2 files changed, 30 deletions(-) diff --git a/web/src/support.ts b/web/src/support.ts index 95ee697628ab..96cb1923b158 100644 --- a/web/src/support.ts +++ b/web/src/support.ts @@ -141,7 +141,6 @@ export class LinearCongruentialGenerator { ); } this.rand_state = state; - this.checkRandState(); } /** diff --git a/web/tests/node/test_random_generator.js b/web/tests/node/test_random_generator.js index 2a7c82e5d195..0e8342202e30 100644 --- a/web/tests/node/test_random_generator.js +++ b/web/tests/node/test_random_generator.js @@ -18,23 +18,6 @@ */ const tvmjs = require("../../dist"); -function createRNGOnlyInstance() { - const tvm = Object.create(tvmjs.Instance.prototype); - tvm.rng = new tvmjs.LinearCongruentialGenerator(); - tvm.empty = function () { - return { - copyFrom(input) { - return { - toArray() { - return input; - }, - }; - }, - }; - }; - return tvm; -} - test("Test coverage of [0,100] inclusive", () => { const covered = Array(100); const rng = new tvmjs.LinearCongruentialGenerator(); @@ -83,18 +66,6 @@ test("Restoring RNG state reproduces next random floats", () => { expect(restored).toEqual(expected); }); -test("Instance RNG state reproduces next uniform samples", () => { - const tvm = createRNGOnlyInstance(); - tvm.setSeed(123); - const state = tvm.getRNGState(); - const expected = Array.from(tvm.uniform([16], -1.0, 1.0, {}).toArray()); - - tvm.uniform([16], -1.0, 1.0, {}); - tvm.setRNGState(state); - const restored = Array.from(tvm.uniform([16], -1.0, 1.0, {}).toArray()); - expect(restored).toEqual(expected); -}); - test("Test two RNGs with different seeds generate different results", () => { const rng1 = new tvmjs.LinearCongruentialGenerator(); const rng2 = new tvmjs.LinearCongruentialGenerator();