diff --git a/.github/workflows/checks.yml b/.github/workflows/checks.yml index 86bd791c2f..b34c9f9d01 100644 --- a/.github/workflows/checks.yml +++ b/.github/workflows/checks.yml @@ -258,6 +258,11 @@ jobs: - name: Execute tests env: TEST_TYPE: ${{ matrix.test_type }} + O1JS_EXPERIMENTAL_MONTGOMERY_MSM: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} + O1JS_EXPERIMENTAL_MONTGOMERY_COMMIT_MSM: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} + O1JS_EXPERIMENTAL_MONTGOMERY_PROVER_MSM: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} + O1JS_EXPERIMENTAL_MONTGOMERY_PROVER_BATCH_MSM: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} + O1JS_MONTGOMERY_MSM_TRACE: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} run: sh run-ci-tests.sh - name: Add to job summary if: always() @@ -444,6 +449,11 @@ jobs: env: TEST_TYPE: ${{ matrix.test_type }} O1JS_BACKEND: native + O1JS_EXPERIMENTAL_MONTGOMERY_MSM: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} + O1JS_EXPERIMENTAL_MONTGOMERY_COMMIT_MSM: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} + O1JS_EXPERIMENTAL_MONTGOMERY_PROVER_MSM: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} + O1JS_EXPERIMENTAL_MONTGOMERY_PROVER_BATCH_MSM: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} + O1JS_MONTGOMERY_MSM_TRACE: ${{ matrix.test_type == 'Performance Regression' && '1' || '0' }} run: sh run-ci-tests.sh - name: Add to job summary if: always() diff --git a/benchmark/benchmarks/transaction.ts b/benchmark/benchmarks/transaction.ts index d6004a5efd..741f8c07c6 100644 --- a/benchmark/benchmarks/transaction.ts +++ b/benchmark/benchmarks/transaction.ts @@ -87,6 +87,14 @@ const TxnBenchmarks = benchmark( await transaction.send(); toc(); }, - // two warmups to ensure full caching - { numberOfWarmups: 2, numberOfRuns: 5 } + // two warmups to ensure full caching by default + { + numberOfWarmups: readBenchmarkCount('O1JS_TXN_BENCH_WARMUPS', 2), + numberOfRuns: readBenchmarkCount('O1JS_TXN_BENCH_RUNS', 5), + } ); + +function readBenchmarkCount(env: string, fallback: number) { + let value = Number(process.env[env] ?? fallback); + return Number.isInteger(value) && value >= 0 ? value : fallback; +} diff --git a/benchmark/runners/srs-msm.ts b/benchmark/runners/srs-msm.ts new file mode 100644 index 0000000000..bf17f0afee --- /dev/null +++ b/benchmark/runners/srs-msm.ts @@ -0,0 +1,183 @@ +/** + * SRS MSM benchmark for the experimental montgomery backend. + * + * Run with: + * ``` + * ./run benchmark/runners/srs-msm.ts + * ``` + * + * Optional knobs: + * - O1JS_MSM_BENCH_SIZES=14,16,18 + * - O1JS_MSM_BENCH_THREADS=4 + * - O1JS_MSM_BENCH_WHOLE_DOMAIN=1 + */ + +import type { MlArray } from '../../src/lib/ml/base.js'; +import type { OrInfinity } from '../../src/bindings/crypto/bindings/curve.js'; +import type { PolyComm } from '../../src/bindings/crypto/bindings/kimchi-types.js'; + +type FieldName = 'fp' | 'fq'; + +const fields: FieldName[] = ['fp', 'fq']; +const logs = parseLogSizes(process.env.O1JS_MSM_BENCH_SIZES ?? '14,16,18'); +const threads = Number(process.env.O1JS_MSM_BENCH_THREADS ?? 0); +const includeWholeDomain = process.env.O1JS_MSM_BENCH_WHOLE_DOMAIN === '1'; +const oldMontgomeryFlag = process.env.O1JS_EXPERIMENTAL_MONTGOMERY_MSM; +const bindings = await importInternal('../../bindings.js', '../../src/bindings.js'); +const { getRustConversion } = await importInternal( + '../../bindings/crypto/bindings.js', + '../../src/bindings/crypto/bindings.js' +); +const { srs: createSrsBindings } = await importInternal( + '../../bindings/crypto/bindings/srs.js', + '../../src/bindings/crypto/bindings/srs.js' +); +const { computeMontgomeryLagrangeCommitment, initializeMontgomeryMsm } = + await importInternal( + '../../bindings/crypto/bindings/montgomery-msm.js', + '../../src/bindings/crypto/bindings/montgomery-msm.js' + ); + +await bindings.initializeBindings(); +let wasm = bindings.wasm; +let conversion = getRustConversion(wasm); +let srsBindings = createSrsBindings(wasm, conversion as any); + +let initialized = await measure('montgomery cold start', () => + initializeMontgomeryMsm({ force: true, threads }) +); +if (initialized.value !== true) { + throw Error( + 'montgomery backend did not initialize. Install optional dependency `montgomery` and run on Node 24+ or a supported browser runtime.' + ); +} + +await bindings.withThreadPool(async () => { + for (let field of fields) { + for (let logSize of logs) { + let domainSize = 1 << logSize; + let srs = ( + await measure(`${field} 2^${logSize} kimchi srs create`, () => + Promise.resolve(srsBindings[field].create(domainSize)) + ) + ).value; + + let indices = [0, domainSize >> 1, domainSize - 1]; + let getSrsPoints = () => + (conversion as any)[field].pointsFromRust( + wasm[`caml_${field}_srs_get`](srs) + ) as MlArray; + + for (let index of indices) { + process.env.O1JS_EXPERIMENTAL_MONTGOMERY_MSM = ''; + let kimchi = await measure(`${field} 2^${logSize} kimchi lagrange[${index}]`, () => + Promise.resolve(srsBindings[field].lagrangeCommitment(srs, domainSize, index)) + ); + + let first = await measure(`${field} 2^${logSize} montgomery first[${index}]`, () => + computeMontgomeryLagrangeCommitment({ + field, + srs, + domainSize, + index, + getSrsPoints, + options: { force: true, threads }, + }) + ); + assertSamePolyComm( + first.value as PolyComm | undefined, + kimchi.value, + `${field} 2^${logSize} first[${index}]` + ); + + let warm = await measure(`${field} 2^${logSize} montgomery warm[${index}]`, () => + computeMontgomeryLagrangeCommitment({ + field, + srs, + domainSize, + index, + getSrsPoints, + options: { force: true, threads }, + }) + ); + assertSamePolyComm( + warm.value as PolyComm | undefined, + kimchi.value, + `${field} 2^${logSize} warm[${index}]` + ); + } + + if (includeWholeDomain) { + process.env.O1JS_EXPERIMENTAL_MONTGOMERY_MSM = ''; + await measure(`${field} 2^${logSize} kimchi whole-domain`, () => + Promise.resolve(srsBindings[field].lagrangeCommitmentsWholeDomain(srs, domainSize)) + ); + } + } + } +}); + +if (oldMontgomeryFlag === undefined) { + delete process.env.O1JS_EXPERIMENTAL_MONTGOMERY_MSM; +} else { + process.env.O1JS_EXPERIMENTAL_MONTGOMERY_MSM = oldMontgomeryFlag; +} + +async function measure(label: string, run: () => Promise) { + let start = performance.now(); + let value = await run(); + let ms = performance.now() - start; + console.log(`${label}: ${ms.toFixed(3)}ms`); + return { value, ms }; +} + +function parseLogSizes(input: string) { + return input + .split(',') + .map((x) => Number(x.trim())) + .filter((x) => Number.isInteger(x) && x > 0); +} + +async function importInternal(distSpecifier: string, sourceSpecifier: string): Promise { + let specifiers = [distSpecifier, sourceSpecifier]; + for (let specifier of specifiers) { + try { + return (await import(specifier)) as T; + } catch (error) { + if (!isModuleNotFound(error)) throw error; + } + } + throw Error(`could not import ${distSpecifier} or ${sourceSpecifier}`); +} + +function isModuleNotFound(error: unknown) { + if (!(error instanceof Error)) return false; + return ( + 'code' in error && + ((error as { code?: unknown }).code === 'ERR_MODULE_NOT_FOUND' || + (error as { code?: unknown }).code === 'MODULE_NOT_FOUND') + ); +} + +function assertSamePolyComm(actual: PolyComm | undefined, expected: PolyComm, label: string) { + if (actual === undefined) throw Error(`${label}: montgomery returned no commitment`); + if (!polyCommEquals(actual, expected)) throw Error(`${label}: commitment mismatch`); +} + +function polyCommEquals(a: PolyComm, b: PolyComm) { + let aPoints = a[1]; + let bPoints = b[1]; + if (aPoints.length !== bPoints.length) return false; + for (let i = 1; i < aPoints.length; i++) { + let aPoint = aPoints[i]; + let bPoint = bPoints[i]; + if (aPoint === 0 || bPoint === 0) { + if (aPoint !== bPoint) return false; + continue; + } + if (aPoint[1][1][1] !== bPoint[1][1][1] || aPoint[1][2][1] !== bPoint[1][2][1]) { + return false; + } + } + return true; +} diff --git a/benchmark/runners/transaction.ts b/benchmark/runners/transaction.ts new file mode 100644 index 0000000000..a9739fae9c --- /dev/null +++ b/benchmark/runners/transaction.ts @@ -0,0 +1,85 @@ +/** + * Focused transaction benchmark runner for end-to-end compile/prove timing. + * + * Run with: + * ``` + * O1JS_TXN_BENCH_RUNS=1 ./run benchmark/runners/transaction.ts + * ``` + */ + +import { + AccountUpdate, + Field, + Mina, + PrivateKey, + SmartContract, + State, + initializeBindings, + method, + state, +} from 'o1js'; + +class SimpleZkapp extends SmartContract { + @state(Field) x = State(); + + init() { + super.init(); + this.x.set(Field(1)); + } + + @method async noOp() {} +} + +await initializeBindings(); + +let runs = readBenchmarkCount('O1JS_TXN_BENCH_RUNS', 1); +let Local = await Mina.LocalBlockchain({ + proofsEnabled: true, + enforceTransactionLimits: true, +}); +Mina.setActiveInstance(Local); + +let transactionFee = 100_000_000; +let [feePayer] = Local.testAccounts; + +let { verificationKey } = await measure('simple zkapp compile', () => SimpleZkapp.compile()); + +for (let run = 0; run < runs; run++) { + let zkappPrivateKey = PrivateKey.random(); + let zkapp = new SimpleZkapp(zkappPrivateKey.toPublicKey()); + + let deployTransaction = await measure('simple zkapp deploy transaction construction', () => + Mina.transaction({ sender: feePayer, fee: transactionFee }, async () => { + AccountUpdate.fundNewAccount(feePayer); + await zkapp.deploy({ verificationKey }); + }) + ); + await measure('simple zkapp deploy transaction signing', async () => { + deployTransaction.sign([feePayer.key, zkappPrivateKey]); + }); + await measure('simple zkapp deploy transaction local sending', () => deployTransaction.send()); + + let callTransaction = await measure('simple zkapp call transaction construction', () => + Mina.transaction({ sender: feePayer, fee: transactionFee }, async () => { + await zkapp.noOp(); + }) + ); + await measure('simple zkapp call transaction proving', () => callTransaction.prove()); + await measure('simple zkapp call transaction signing', async () => { + callTransaction.sign([feePayer.key]); + }); + await measure('simple zkapp call transaction local sending', () => callTransaction.send()); +} + +async function measure(label: string, run: () => Promise) { + let start = performance.now(); + let value = await run(); + let ms = performance.now() - start; + console.log(`${label}: ${ms.toFixed(3)}ms`); + return value; +} + +function readBenchmarkCount(env: string, fallback: number) { + let value = Number(process.env[env] ?? fallback); + return Number.isInteger(value) && value >= 0 ? value : fallback; +} diff --git a/npmDepsHash b/npmDepsHash index 2e182bf0a6..e43da7cbcf 100644 --- a/npmDepsHash +++ b/npmDepsHash @@ -1 +1 @@ -sha256-pKzMJ78FJBWm6SxOoRmzBVbsUZRUSHgX/oXOZO0Aqtw= +sha256-1fIW5ylFDZFs3PLIZpTYxuWH/meXL95umHri/jqSK0c= diff --git a/package-lock.json b/package-lock.json index c19df83e58..ae8236a0a9 100644 --- a/package-lock.json +++ b/package-lock.json @@ -56,7 +56,8 @@ "node": ">=18.14.0" }, "optionalDependencies": { - "@o1js/native": "2.15.0" + "@o1js/native": "2.15.0", + "montgomery": "git+https://github.com/Trivo25/montgomery.git#florian/sync-9x29" } }, "native/darwin-arm64": { @@ -3407,6 +3408,67 @@ } }, "node_modules/@o1js/native": { + "version": "2.15.0", + "resolved": "https://registry.npmjs.org/@o1js/native/-/native-2.15.0.tgz", + "integrity": "sha512-te/piape+Nn9bgctqZk1ARnBMilKwq+IkWT4Z4jKh3HNuGj/Z5rngVGMRAPzJzxyjQEh7Dm+/BN6M9D6VY0kdw==", + "optional": true, + "optionalDependencies": { + "@o1js/native-darwin-arm64": "2.15.0", + "@o1js/native-darwin-x64": "2.15.0", + "@o1js/native-linux-arm64": "2.15.0", + "@o1js/native-linux-x64": "2.15.0", + "@o1js/native-win32-x64": "2.15.0" + } + }, + "node_modules/@o1js/native-darwin-arm64": { + "version": "2.15.0", + "resolved": "https://registry.npmjs.org/@o1js/native-darwin-arm64/-/native-darwin-arm64-2.15.0.tgz", + "integrity": "sha512-4S1hf6Z6KLI1ITnbecMS6y3q3D+IqIGjF9vC5xSHQETUGPS0IFbWGUwai/52RIfH65eXBK5MXWYwwAlc6rcz5g==", + "cpu": [ + "arm64" + ], + "optional": true, + "os": [ + "darwin" + ] + }, + "node_modules/@o1js/native-darwin-x64": { + "version": "2.15.0", + "resolved": "https://registry.npmjs.org/@o1js/native-darwin-x64/-/native-darwin-x64-2.15.0.tgz", + "integrity": "sha512-zsCvwBJK3j9yDcfpdK0QypaQKTRFnQTmCH989JEYr61yqFQ5Uxz03oEVkCU02p2+w2bZzKWIy/7rTTAVaCiYqg==", + "cpu": [ + "x64" + ], + "optional": true, + "os": [ + "darwin" + ] + }, + "node_modules/@o1js/native-linux-arm64": { + "version": "2.15.0", + "resolved": "https://registry.npmjs.org/@o1js/native-linux-arm64/-/native-linux-arm64-2.15.0.tgz", + "integrity": "sha512-Fa1tWgvrZSCDpJ2c4rHZKnfhxpZc/XkqGNVW5gDocr6FArSpWyEyKn0tPJJKtDUIzkUCozjxnFO8EWws0HJ6MA==", + "cpu": [ + "arm64" + ], + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@o1js/native-linux-x64": { + "version": "2.15.0", + "resolved": "https://registry.npmjs.org/@o1js/native-linux-x64/-/native-linux-x64-2.15.0.tgz", + "integrity": "sha512-D2Z3NBb7mtehYRzPPlCT+m4/ojJYyo0MJc/9f0MtwIXIlTRL1jVBg6qH8TQdIK5j3LYlFuAhjDqaWfjAxQGwMA==", + "cpu": [ + "x64" + ], + "optional": true, + "os": [ + "linux" + ] + }, + "node_modules/@o1js/native/node_modules/@o1js/native-win32-x64": { "optional": true }, "node_modules/@octokit/action": { @@ -5389,6 +5451,27 @@ "url": "https://opencollective.com/express" } }, + "node_modules/ieee754": { + "version": "1.2.1", + "resolved": "https://registry.npmjs.org/ieee754/-/ieee754-1.2.1.tgz", + "integrity": "sha512-dcyqhDvX1C46lXZcVqCpK+FtMRQVdIMN6/Df5js2zouUsqG7I6sFxitIC+7KYK29KdXOLHdu9zL4sFnoVQnqaA==", + "funding": [ + { + "type": "github", + "url": "https://github.com/sponsors/feross" + }, + { + "type": "patreon", + "url": "https://www.patreon.com/feross" + }, + { + "type": "consulting", + "url": "https://feross.org/support" + } + ], + "license": "BSD-3-Clause", + "optional": true + }, "node_modules/ignore": { "version": "5.3.1", "resolved": "https://registry.npmjs.org/ignore/-/ignore-5.3.1.tgz", @@ -7284,6 +7367,15 @@ "ufo": "^1.5.3" } }, + "node_modules/montgomery": { + "version": "0.4.0", + "resolved": "git+https://github.com/Trivo25/montgomery.git#6ef95d1f6ee233f0f730f5d369ba9e2759898ba7", + "license": "Apache-2.0", + "optional": true, + "dependencies": { + "wasmati": "^0.2.2" + } + }, "node_modules/ms": { "version": "2.1.3", "resolved": "https://registry.npmjs.org/ms/-/ms-2.1.3.tgz", @@ -8592,6 +8684,16 @@ "makeerror": "1.0.12" } }, + "node_modules/wasmati": { + "version": "0.2.4", + "resolved": "https://registry.npmjs.org/wasmati/-/wasmati-0.2.4.tgz", + "integrity": "sha512-cJl05zeUhYM5Vx6A8p4z8S9H9HVJPlLK4HO5RFoXswy5ZoWjd2zN0FX1amOs3vkshxbiKFQpJsefkK9OBgyYjQ==", + "license": "MIT", + "optional": true, + "dependencies": { + "ieee754": "^1.2.1" + } + }, "node_modules/which": { "version": "2.0.2", "resolved": "https://registry.npmjs.org/which/-/which-2.0.2.tgz", diff --git a/package.json b/package.json index 20307cfa11..f3236a50be 100644 --- a/package.json +++ b/package.json @@ -137,7 +137,6 @@ }, "dependencies": { "@noble/hashes": "^1.7.1", - "@o1js/native": "file:./native/meta", "blakejs": "1.2.1", "cachedir": "^2.4.0", "libsodium-wrappers-sumo": "^0.7.15", @@ -149,6 +148,7 @@ "native-version": "2.15.0" }, "optionalDependencies": { - "@o1js/native": "2.15.0" + "@o1js/native": "2.15.0", + "montgomery": "git+https://github.com/Trivo25/montgomery.git#florian/sync-9x29" } } diff --git a/run-in-browser.js b/run-in-browser.js index 5e275b606f..0893f5f16a 100755 --- a/run-in-browser.js +++ b/run-in-browser.js @@ -26,6 +26,15 @@ await fs.copyFile(absPath, newPath); await fs.unlink(absPath); console.log(`running in the browser: ${newPath}`); +let browserFlags = []; +if (process.env.O1JS_EXPERIMENTAL_MONTGOMERY_MSM === '1') { + browserFlags.push(`globalThis.O1JS_EXPERIMENTAL_MONTGOMERY_MSM = '1';`); +} +let browserFlagsScript = + browserFlags.length === 0 + ? '' + : ``; + const indexHtml = ` @@ -33,6 +42,7 @@ const indexHtml = ` o1js + ${browserFlagsScript} diff --git a/src/bindings.js b/src/bindings.js index 8a25c51b4a..d8a52abec2 100644 --- a/src/bindings.js +++ b/src/bindings.js @@ -21,11 +21,14 @@ async function initializeBindings() { lockBackend(); + let backend; if (getBackendPreference() === 'native') { - ({ wasm, withThreadPool } = await import('./bindings/js/node/native-backend.js')); + backend = await import('./bindings/js/node/native-backend.js'); } else { - ({ wasm, withThreadPool } = await import('./bindings/js/node/node-backend.js')); + backend = await import('./bindings/js/node/node-backend.js'); + await backend.montgomeryBridgeReady; } + ({ wasm, withThreadPool } = backend); // this dynamic import makes jest respect the import order // otherwise the cjs file gets imported before its implicit esm dependencies and fails diff --git a/src/bindings/crypto/bindings/montgomery-msm.ts b/src/bindings/crypto/bindings/montgomery-msm.ts new file mode 100644 index 0000000000..c8101f0528 --- /dev/null +++ b/src/bindings/crypto/bindings/montgomery-msm.ts @@ -0,0 +1,621 @@ +import { MlArray } from '../../../lib/ml/base.js'; +import { getBackendPreference } from '../../../lib/backend.js'; +import { Fp, Fq, type FiniteField } from '../finite-field.js'; +import { OrInfinity } from './curve.js'; +import { Field } from './field.js'; +import { PolyComm } from './kimchi-types.js'; + +export { + computeMontgomeryCommitEvaluations, + computeMontgomeryLagrangeCommitment, + getCachedMontgomeryCommitEvaluations, + getCachedMontgomeryLagrangeCommitment, + getCachedMontgomeryLagrangeCommitmentsWholeDomain, + initializeMontgomeryMsm, + isMontgomeryMsmEnabled, + msm, + msmUnsafe, + precomputeMontgomeryCommitEvaluations, + precomputeMontgomeryLagrangeCommitment, + warmupMontgomeryCommitEvaluations, + warmupMontgomeryLagrangeCommitment, + warmupMontgomeryLagrangeCommitmentsWholeDomain, +}; + +type FieldName = 'fp' | 'fq'; +type AffinePoint = { x: bigint; y: bigint; isZero: boolean }; +type MontgomeryModule = { + Pallas(): Promise; + Vesta(): Promise; + startThreads?(threads?: number): Promise; +}; +type MontgomeryCurve = { + Field: { + local: { getPointer(size?: number): number; getPointers(n: number): number[] }; + getPointer(size?: number): number; + }; + Scalar: { + sizeField: number; + writeBigint(pointer: number, value: bigint): void; + }; + Affine: { + size: number; + writeBigints(pointPtr: number, points: AffinePoint[]): void; + toBigint(point: number): AffinePoint; + }; + Projective: { toAffine(scratch: number[], affine: number, point: number): void }; + Parallel: { + getPointer(size: number): Promise; + getScalarPointer(size: number): Promise; + msm( + scalarPtr: number, + pointPtr: number, + n: number, + verboseTiming?: boolean, + options?: { c?: number; useSafeAdditions?: boolean } + ): Promise<{ result: number; log: unknown[][] }>; + msmUnsafe( + scalarPtr: number, + pointPtr: number, + n: number, + verboseTiming?: boolean, + options?: { c?: number; c0?: number } + ): Promise<{ result: number; log: unknown[][] }>; + }; +}; + +type MontgomeryOptions = { + force?: boolean; + threads?: number; +}; + +type SrsPointCache = { pointPtr: number; length: number }; +type LagrangeCache = { + commitments: Map; + wholeDomains: Map>; + pending: Map>; +}; +type CommitEvaluationsCache = { + commitments: WeakMap>; + pending: WeakMap>>; +}; + +const srsPoints = new WeakMap>>(); +const lagrangeCache = new WeakMap>(); +const commitEvaluationsCache = new WeakMap>(); +const scalarPointers = new Map>(); + +let modulePromise: Promise | undefined; +let montgomeryModule: MontgomeryModule | undefined; +let threadPromise: Promise | undefined; +let curves: Partial>> = {}; + +function isMontgomeryMsmEnabled(options?: MontgomeryOptions) { + if (options?.force === true) return true; + if (getBackendPreference() !== 'wasm') return false; + let globalFlag = (globalThis as { O1JS_EXPERIMENTAL_MONTGOMERY_MSM?: unknown }) + .O1JS_EXPERIMENTAL_MONTGOMERY_MSM; + let processFlag = + typeof process !== 'undefined' ? process.env.O1JS_EXPERIMENTAL_MONTGOMERY_MSM : undefined; + return globalFlag === '1' || globalFlag === true || processFlag === '1'; +} + +async function initializeMontgomeryMsm(options?: MontgomeryOptions) { + let module = await loadMontgomery(options); + if (module === undefined) return false; + let [fp, fq] = await Promise.all([curveForField('fp', options), curveForField('fq', options)]); + return fp !== undefined && fq !== undefined; +} + +async function loadMontgomery(options?: MontgomeryOptions) { + if (!isMontgomeryMsmEnabled(options) || !runtimeSupportsMontgomery()) return undefined; + if (montgomeryModule !== undefined) return montgomeryModule; + modulePromise ??= importOptionalMontgomery().then(async (module) => { + if (module === undefined) return undefined; + let threads = options?.threads ?? configuredThreadCount(); + if (threads > 0 && typeof module.startThreads === 'function') { + threadPromise ??= module.startThreads(threads); + await threadPromise; + } + montgomeryModule = module; + return module; + }); + return modulePromise; +} + +async function importOptionalMontgomery(): Promise { + try { + let specifier = 'montgomery'; + return (await import(specifier)) as MontgomeryModule; + } catch { + return undefined; + } +} + +async function curveForField(field: FieldName, options?: MontgomeryOptions) { + let cached = curves[field]; + if (cached !== undefined) return cached; + let module = await loadMontgomery(options); + if (module === undefined) return undefined; + let curve = field === 'fp' ? module.Vesta() : module.Pallas(); + curves[field] = curve; + return curve; +} + +function runtimeSupportsMontgomery() { + if ( + typeof SharedArrayBuffer === 'undefined' || + typeof (Atomics as any).waitAsync !== 'function' + ) { + return false; + } + if (typeof process === 'undefined' || process.release?.name !== 'node') return true; + let major = Number(process.versions.node.split('.')[0]); + return major >= 24; +} + +function configuredThreadCount() { + if (typeof process === 'undefined') return 0; + let value = Number(process.env.O1JS_MONTGOMERY_THREADS ?? 0); + return Number.isInteger(value) && value > 0 ? value : 0; +} + +async function msm( + field: FieldName, + points: OrInfinity[], + scalars: bigint[], + options?: MontgomeryOptions +) { + return msmImpl(field, points, scalars, false, options); +} + +async function msmUnsafe( + field: FieldName, + points: OrInfinity[], + scalars: bigint[], + options?: MontgomeryOptions +) { + return msmImpl(field, points, scalars, true, options); +} + +async function msmImpl( + field: FieldName, + points: OrInfinity[], + scalars: bigint[], + unsafe: boolean, + options?: MontgomeryOptions +): Promise { + if (points.length !== scalars.length) throw Error('montgomery MSM input length mismatch'); + let curve = await curveForField(field, options); + if (curve === undefined) return undefined; + let pointPtr = await writeMontgomeryPoints(curve, points.map(orInfinityToMontgomery)); + let scalarPtr = await writeMontgomeryScalars(curve, scalars); + let { result } = unsafe + ? await curve.Parallel.msmUnsafe(scalarPtr, pointPtr, scalars.length) + : await curve.Parallel.msm(scalarPtr, pointPtr, scalars.length); + return polyCommFromMontgomeryResult(curve, result); +} + +function getCachedMontgomeryCommitEvaluations( + field: FieldName, + srs: object, + domainSize: number, + evaluations: MlArray +) { + if (!isMontgomeryMsmEnabled()) return undefined; + return getCommitEvaluationsValueMap( + getCommitEvaluationsCache(srs, field).commitments, + evaluations + ).get(domainSize); +} + +function warmupMontgomeryCommitEvaluations(input: { + field: FieldName; + srs: object; + domainSize: number; + evaluations: MlArray; + getSrsPoints(): MlArray; + expected?: PolyComm; +}) { + void precomputeMontgomeryCommitEvaluations(input).catch(() => undefined); +} + +async function precomputeMontgomeryCommitEvaluations(input: { + field: FieldName; + srs: object; + domainSize: number; + evaluations: MlArray; + getSrsPoints(): MlArray; + expected?: PolyComm; + options?: MontgomeryOptions; +}) { + if (!isMontgomeryMsmEnabled(input.options)) return false; + let cache = getCommitEvaluationsCache(input.srs, input.field); + let commitmentMap = getCommitEvaluationsValueMap(cache.commitments, input.evaluations); + if (commitmentMap.has(input.domainSize)) return true; + + let pendingMap = getCommitEvaluationsValueMap(cache.pending, input.evaluations); + let pending = pendingMap.get(input.domainSize); + if (pending !== undefined) return pending; + + let promise = computeMontgomeryCommitEvaluations(input) + .then((commitment) => { + if (commitment === undefined) return false; + if (input.expected !== undefined && !polyCommEquals(commitment, input.expected)) return false; + getCommitEvaluationsValueMap( + getCommitEvaluationsCache(input.srs, input.field).commitments, + input.evaluations + ).set(input.domainSize, commitment); + return true; + }) + .catch(() => false) + .finally(() => { + let pendingMap = getCommitEvaluationsValueMap( + getCommitEvaluationsCache(input.srs, input.field).pending, + input.evaluations + ); + if (pendingMap.get(input.domainSize) === promise) pendingMap.delete(input.domainSize); + }); + pendingMap.set(input.domainSize, promise); + return promise; +} + +async function computeMontgomeryCommitEvaluations(input: { + field: FieldName; + srs: object; + domainSize: number; + evaluations: MlArray; + getSrsPoints(): MlArray; + options?: MontgomeryOptions; +}) { + let curve = await curveForField(input.field, input.options); + if (curve === undefined) return undefined; + let points = await getMontgomerySrsPoints(input.field, input.srs, input.getSrsPoints, curve); + if (points.length < input.domainSize) return undefined; + let scalars = interpolateEvaluations( + input.field, + MlArray.mapFrom(input.evaluations, ([, x]) => x) + ); + if (scalars.length !== input.domainSize) return undefined; + let scalarPtr = await writeMontgomeryScalars(curve, scalars); + let { result } = await curve.Parallel.msm(scalarPtr, points.pointPtr, input.domainSize); + return polyCommFromMontgomeryResult(curve, result); +} + +function getCachedMontgomeryLagrangeCommitment( + field: FieldName, + srs: object, + domainSize: number, + index: number +) { + if (!isMontgomeryMsmEnabled()) return undefined; + return getLagrangeCache(srs, field).commitments.get(lagrangeKey(domainSize, index)); +} + +function warmupMontgomeryLagrangeCommitment(input: { + field: FieldName; + srs: object; + domainSize: number; + index: number; + getSrsPoints(): MlArray; + expected?: PolyComm; +}) { + void precomputeMontgomeryLagrangeCommitment(input).catch(() => undefined); +} + +async function precomputeMontgomeryLagrangeCommitment(input: { + field: FieldName; + srs: object; + domainSize: number; + index: number; + getSrsPoints(): MlArray; + expected?: PolyComm; + options?: MontgomeryOptions; +}) { + if (!isMontgomeryMsmEnabled(input.options)) return false; + let cache = getLagrangeCache(input.srs, input.field); + let key = lagrangeKey(input.domainSize, input.index); + if (cache.commitments.has(key)) return true; + let pending = cache.pending.get(key); + if (pending !== undefined) return pending; + + let promise = computeMontgomeryLagrangeCommitment(input) + .then((commitment) => { + if (commitment === undefined) return false; + if (input.expected !== undefined && !polyCommEquals(commitment, input.expected)) return false; + cache.commitments.set(key, commitment); + return true; + }) + .catch(() => false) + .finally(() => { + if (cache.pending.get(key) === promise) cache.pending.delete(key); + }); + cache.pending.set(key, promise); + return promise; +} + +async function computeMontgomeryLagrangeCommitment(input: { + field: FieldName; + srs: object; + domainSize: number; + index: number; + getSrsPoints(): MlArray; + options?: MontgomeryOptions; +}) { + let curve = await curveForField(input.field, input.options); + if (curve === undefined) return undefined; + let points = await getMontgomerySrsPoints(input.field, input.srs, input.getSrsPoints, curve); + if (points.length < input.domainSize) return undefined; + let scalars = await getLagrangeScalars(input.field, input.domainSize, input.index, curve); + let { result } = await curve.Parallel.msmUnsafe( + scalars.scalarPtr, + points.pointPtr, + input.domainSize + ); + return polyCommFromMontgomeryResult(curve, result); +} + +function getCachedMontgomeryLagrangeCommitmentsWholeDomain( + field: FieldName, + srs: object, + domainSize: number +) { + if (!isMontgomeryMsmEnabled()) return undefined; + return getLagrangeCache(srs, field).wholeDomains.get(domainSize); +} + +function warmupMontgomeryLagrangeCommitmentsWholeDomain(input: { + field: FieldName; + srs: object; + domainSize: number; + getSrsPoints(): MlArray; + expected?: MlArray; +}) { + if (!isMontgomeryMsmEnabled()) return; + let max = configuredWholeDomainWarmupMax(); + if (max <= 0 || input.domainSize > max) return; + let cache = getLagrangeCache(input.srs, input.field); + let key = `whole:${input.domainSize}`; + if (cache.wholeDomains.has(input.domainSize) || cache.pending.has(key)) return; + let promise = Promise.all( + Array.from({ length: input.domainSize }, (_, index) => + computeMontgomeryLagrangeCommitment({ ...input, index }) + ) + ) + .then((commitments) => { + if (commitments.some((commitment) => commitment === undefined)) return false; + let mlCommitments = [0, ...commitments] as MlArray; + if (input.expected !== undefined && !polyCommsEqual(mlCommitments, input.expected)) { + return false; + } + cache.wholeDomains.set(input.domainSize, mlCommitments); + return true; + }) + .catch(() => false) + .finally(() => { + if (cache.pending.get(key) === promise) cache.pending.delete(key); + }); + cache.pending.set(key, promise); + void promise; +} + +function configuredWholeDomainWarmupMax() { + if (typeof process === 'undefined') return 0; + let value = Number(process.env.O1JS_MONTGOMERY_WHOLE_DOMAIN_MAX ?? 0); + return Number.isInteger(value) && value > 0 ? value : 0; +} + +async function getMontgomerySrsPoints( + field: FieldName, + srs: object, + getSrsPoints: () => MlArray, + curve: MontgomeryCurve +) { + let fieldCache = srsPoints.get(srs); + if (fieldCache === undefined) srsPoints.set(srs, (fieldCache = new Map())); + let cached = fieldCache.get(field); + if (cached !== undefined) return cached; + let promise = Promise.resolve().then(async () => { + let [, _h, ...gs] = getSrsPoints(); + let points = gs.map(orInfinityToMontgomery); + return { pointPtr: await writeMontgomeryPoints(curve, points), length: points.length }; + }); + fieldCache.set(field, promise); + return promise; +} + +async function getLagrangeScalars( + field: FieldName, + domainSize: number, + index: number, + curve: MontgomeryCurve +) { + let key = `${field}:${domainSize}:${index}`; + let cached = scalarPointers.get(key); + if (cached !== undefined) return cached; + let promise = Promise.resolve().then(async () => { + let scalars = lagrangeScalars(field, domainSize, index); + return { scalarPtr: await writeMontgomeryScalars(curve, scalars), length: scalars.length }; + }); + scalarPointers.set(key, promise); + return promise; +} + +async function writeMontgomeryPoints(curve: MontgomeryCurve, points: AffinePoint[]) { + let pointPtr = await curve.Parallel.getPointer(points.length * curve.Affine.size); + curve.Affine.writeBigints(pointPtr, points); + return pointPtr; +} + +async function writeMontgomeryScalars(curve: MontgomeryCurve, scalars: bigint[]) { + let scalarPtr = await curve.Parallel.getScalarPointer(scalars.length * curve.Scalar.sizeField); + for (let i = 0, ptr = scalarPtr; i < scalars.length; i++, ptr += curve.Scalar.sizeField) { + curve.Scalar.writeBigint(ptr, scalars[i]); + } + return scalarPtr; +} + +function lagrangeScalars(field: FieldName, domainSize: number, index: number) { + if (!Number.isInteger(domainSize) || domainSize <= 0 || (domainSize & (domainSize - 1)) !== 0) { + throw Error(`expected power-of-two domain size, got ${domainSize}`); + } + if (!Number.isInteger(index) || index < 0 || index >= domainSize) { + throw Error(`lagrange index ${index} out of range for domain ${domainSize}`); + } + let Field = field === 'fp' ? Fp : Fq; + let omega = domainGenerator(Field, Math.log2(domainSize)); + let omegaInverse = Field.inverse(omega); + let domainSizeInverse = Field.inverse(BigInt(domainSize)); + if (omegaInverse === undefined || domainSizeInverse === undefined) { + throw Error('invalid lagrange domain'); + } + let step = Field.power(omegaInverse, BigInt(index)); + let scalar = domainSizeInverse; + let scalars = Array(domainSize); + for (let i = 0; i < domainSize; i++) { + scalars[i] = scalar; + scalar = Field.mul(scalar, step); + } + return scalars; +} + +function interpolateEvaluations(field: FieldName, evaluations: bigint[]) { + let n = evaluations.length; + if (!Number.isInteger(n) || n <= 0 || (n & (n - 1)) !== 0) { + throw Error(`expected power-of-two evaluation length, got ${n}`); + } + let Field = field === 'fp' ? Fp : Fq; + let omega = domainGenerator(Field, Math.log2(n)); + let omegaInverse = Field.inverse(omega); + let nInverse = Field.inverse(BigInt(n)); + if (omegaInverse === undefined || nInverse === undefined) { + throw Error('invalid evaluation domain'); + } + let coefficients = fft(Field, evaluations, omegaInverse); + for (let i = 0; i < coefficients.length; i++) { + coefficients[i] = Field.mul(coefficients[i], nInverse); + } + return coefficients; +} + +function fft(Field: FiniteField, input: bigint[], omega: bigint) { + let n = input.length; + let values = input.slice(); + bitReverseInPlace(values); + for (let length = 2; length <= n; length <<= 1) { + let step = Field.power(omega, BigInt(n / length)); + for (let offset = 0; offset < n; offset += length) { + let twiddle = 1n; + let half = length >> 1; + for (let j = 0; j < half; j++) { + let even = values[offset + j]; + let odd = Field.mul(values[offset + j + half], twiddle); + values[offset + j] = Field.add(even, odd); + values[offset + j + half] = Field.sub(even, odd); + twiddle = Field.mul(twiddle, step); + } + } + } + return values; +} + +function bitReverseInPlace(values: bigint[]) { + let n = values.length; + for (let i = 1, j = 0; i < n; i++) { + let bit = n >> 1; + for (; (j & bit) !== 0; bit >>= 1) j ^= bit; + j ^= bit; + if (i < j) { + let tmp = values[i]; + values[i] = values[j]; + values[j] = tmp; + } + } +} + +function domainGenerator(Field: FiniteField, logSize: number) { + if (logSize > Number(Field.M) || logSize < 0) { + throw Error(`log2 size of evaluation domain must be in [0, ${Field.M}], got ${logSize}`); + } + let generator = Field.twoadicRoot; + for (let j = Number(Field.M); j > logSize; j--) { + generator = Field.square(generator); + } + return generator; +} + +function orInfinityToMontgomery(point: OrInfinity): AffinePoint { + if (point === 0) throw Error('montgomery MSM fast path does not accept infinity points'); + return { x: point[1][1][1], y: point[1][2][1], isZero: false }; +} + +function polyCommFromMontgomeryResult(curve: MontgomeryCurve, result: number): PolyComm { + let scratch = curve.Field.local.getPointers(5); + let affine = curve.Field.local.getPointer(curve.Affine.size); + curve.Projective.toAffine(scratch, affine, result); + let point = curve.Affine.toBigint(affine); + let mlPoint: OrInfinity = point.isZero ? 0 : [0, [0, [0, point.x], [0, point.y]]]; + return [0, [0, mlPoint]]; +} + +function getLagrangeCache(srs: object, field: FieldName) { + let srsCache = lagrangeCache.get(srs); + if (srsCache === undefined) lagrangeCache.set(srs, (srsCache = new Map())); + let fieldCache = srsCache.get(field); + if (fieldCache === undefined) { + fieldCache = { commitments: new Map(), wholeDomains: new Map(), pending: new Map() }; + srsCache.set(field, fieldCache); + } + return fieldCache; +} + +function getCommitEvaluationsCache(srs: object, field: FieldName) { + let srsCache = commitEvaluationsCache.get(srs); + if (srsCache === undefined) commitEvaluationsCache.set(srs, (srsCache = new Map())); + let fieldCache = srsCache.get(field); + if (fieldCache === undefined) { + fieldCache = { commitments: new WeakMap(), pending: new WeakMap() }; + srsCache.set(field, fieldCache); + } + return fieldCache; +} + +function getCommitEvaluationsValueMap( + cache: WeakMap>, + evaluations: MlArray +) { + let key = evaluations as unknown as object; + let map = cache.get(key); + if (map === undefined) { + map = new Map(); + cache.set(key, map); + } + return map; +} + +function lagrangeKey(domainSize: number, index: number) { + return `${domainSize}:${index}`; +} + +function polyCommsEqual(a: MlArray, b: MlArray) { + if (a.length !== b.length) return false; + for (let i = 1; i < a.length; i++) { + if (!polyCommEquals(a[i] as PolyComm, b[i] as PolyComm)) return false; + } + return true; +} + +function polyCommEquals(a: PolyComm, b: PolyComm) { + let aPoints = a[1]; + let bPoints = b[1]; + if (aPoints.length !== bPoints.length) return false; + for (let i = 1; i < aPoints.length; i++) { + if (!orInfinityEquals(aPoints[i], bPoints[i])) return false; + } + return true; +} + +function orInfinityEquals(a: OrInfinity, b: OrInfinity) { + if (a === 0 || b === 0) return a === b; + return a[1][1][1] === b[1][1][1] && a[1][2][1] === b[1][2][1]; +} diff --git a/src/bindings/crypto/bindings/srs.ts b/src/bindings/crypto/bindings/srs.ts index e6c5e67cf4..d0f870664e 100644 --- a/src/bindings/crypto/bindings/srs.ts +++ b/src/bindings/crypto/bindings/srs.ts @@ -1,17 +1,26 @@ -import type { Wasm, RustConversion } from '../bindings.js'; -import { type WasmFpSrs, type WasmFqSrs } from '../../compiled/node_bindings/kimchi_wasm.cjs'; -import { PolyComm } from './kimchi-types.js'; -import { srsCache as cache } from '../cache.js'; +import { MlArray } from '../../../lib/ml/base.js'; import { - type CacheHeader, - type Cache, + readCache, withVersion, writeCache, - readCache, + type Cache, + type CacheHeader, } from '../../../lib/proof-system/cache.js'; import { assert } from '../../../lib/util/errors.js'; -import { MlArray } from '../../../lib/ml/base.js'; +import { type WasmFpSrs, type WasmFqSrs } from '../../compiled/node_bindings/kimchi_wasm.cjs'; +import type { RustConversion, Wasm } from '../bindings.js'; +import { srsCache as cache } from '../cache.js'; import { OrInfinity, OrInfinityJson } from './curve.js'; +import { Field } from './field.js'; +import { PolyComm } from './kimchi-types.js'; +import { + getCachedMontgomeryCommitEvaluations, + getCachedMontgomeryLagrangeCommitment, + getCachedMontgomeryLagrangeCommitmentsWholeDomain, + warmupMontgomeryCommitEvaluations, + warmupMontgomeryLagrangeCommitment, + warmupMontgomeryLagrangeCommitmentsWholeDomain, +} from './montgomery-msm.js'; export { srs }; @@ -141,7 +150,22 @@ function srsPerField(f: 'fp' | 'fq', wasm: Wasm, conversion: RustConversion<'was /** * returns ith Lagrange basis commitment for a given domain size */ - lagrangeCommitment(srs: WasmSrs, domainSize: number, i: number): PolyComm { + lagrangeCommitment( + srs: WasmSrs, + domainSize: number, + i: number, + options?: { skipMontgomeryCache?: boolean } + ): PolyComm { + if (options?.skipMontgomeryCache !== true) { + let cachedMontgomeryCommitment = getCachedMontgomeryLagrangeCommitment( + f, + srs, + domainSize, + i + ); + if (cachedMontgomeryCommitment !== undefined) return cachedMontgomeryCommitment; + } + // happy, fast case: if basis is already stored on the srs, return the ith commitment let commitment = maybeLagrangeCommitment(srs, domainSize, i); @@ -208,13 +232,29 @@ function srsPerField(f: 'fp' | 'fq', wasm: Wasm, conversion: RustConversion<'was writeCache(cache, header, bytes); } } - return conversion[f].polyCommFromRust(commitment); + let mlCommitment = conversion[f].polyCommFromRust(commitment); + warmupMontgomeryLagrangeCommitment({ + field: f, + srs, + domainSize, + index: i, + getSrsPoints: () => conversion[f].pointsFromRust(getSrs(srs)), + expected: mlCommitment, + }); + return mlCommitment; }, /** * Returns the Lagrange basis commitments for the whole domain */ lagrangeCommitmentsWholeDomain(srs: WasmSrs, domainSize: number) { + let cachedMontgomeryCommitments = getCachedMontgomeryLagrangeCommitmentsWholeDomain( + f, + srs, + domainSize + ); + if (cachedMontgomeryCommitments !== undefined) return cachedMontgomeryCommitments; + // instead of getting the entire commitment directly (which works for nodejs/servers), we get a pointer to the commitment // and then read the commitment from the pointer // this is because the web worker implementation currently does not support returning UintXArray's directly @@ -224,6 +264,13 @@ function srsPerField(f: 'fp' | 'fq', wasm: Wasm, conversion: RustConversion<'was let ptr = lagrangeCommitmentsWholeDomainPtr(srs, domainSize); let wasmComms = getCommitmentsWholeDomainByPtr(ptr); let mlComms = conversion[f].polyCommsFromRust(wasmComms); + warmupMontgomeryLagrangeCommitmentsWholeDomain({ + field: f, + srs, + domainSize, + getSrsPoints: () => conversion[f].pointsFromRust(getSrs(srs)), + expected: mlComms, + }); return mlComms; }, @@ -232,7 +279,27 @@ function srsPerField(f: 'fp' | 'fq', wasm: Wasm, conversion: RustConversion<'was */ addLagrangeBasis(srs: WasmSrs, logSize: number) { // this ensures that basis is stored on the srs, no need to duplicate caching logic - this.lagrangeCommitment(srs, 1 << logSize, 0); + this.lagrangeCommitment(srs, 1 << logSize, 0, { skipMontgomeryCache: true }); + }, + + commitEvaluationsCached(srs: WasmSrs, domainSize: number, evaluations: MlArray) { + return getCachedMontgomeryCommitEvaluations(f, srs, domainSize, evaluations); + }, + + warmupCommitEvaluations(input: { + srs: WasmSrs; + domainSize: number; + evaluations: MlArray; + expected?: PolyComm; + }) { + warmupMontgomeryCommitEvaluations({ + field: f, + srs: input.srs, + domainSize: input.domainSize, + evaluations: input.evaluations, + getSrsPoints: () => conversion[f].pointsFromRust(getSrs(input.srs)), + expected: input.expected, + }); }, }; } diff --git a/src/bindings/js/montgomery-msm-bridge.js b/src/bindings/js/montgomery-msm-bridge.js new file mode 100644 index 0000000000..d00bf9a073 --- /dev/null +++ b/src/bindings/js/montgomery-msm-bridge.js @@ -0,0 +1,223 @@ +let bridgeReady; +let curves; + +export { + installMontgomeryMsmBridge, + isMontgomeryProverBatchMsmEnabled, + isMontgomeryProverMsmEnabled, +}; + +function installMontgomeryMsmBridge() { + if (bridgeReady !== undefined) return bridgeReady; + globalThis.o1jsMontgomerySrsMsmEnabled = isMontgomeryBridgeReady; + globalThis.o1jsMontgomeryCommitMsmEnabled = isMontgomeryCommitBridgeReady; + globalThis.o1jsMontgomeryProverMsmEnabled = isMontgomeryProverBridgeReady; + globalThis.o1jsMontgomeryProverBatchMsmEnabled = isMontgomeryProverBatchBridgeReady; + globalThis.o1jsMontgomerySrsMsm = montgomerySrsMsm; + globalThis.o1jsMontgomerySrsMsmBatch = montgomerySrsMsmBatch; + bridgeReady = isMontgomeryBridgeEnabled() ? initializeCurves().catch(() => undefined) : undefined; + return bridgeReady; +} + +async function initializeCurves() { + let specifier = 'montgomery'; + let montgomery = await import(specifier); + let [pallas, vesta] = await Promise.all([montgomery.Pallas(), montgomery.Vesta()]); + curves = { pallas, vesta }; +} + +function montgomerySrsMsm(curveName, pointBytes, scalarBytes) { + if (!isMontgomeryBridgeEnabled() || curves === undefined) return undefined; + let curve = curves[curveName]; + if (curve === undefined) return undefined; + let n = pointBytes.length / 64; + if (typeof curve.Sync?.msmAffineBytesBatch === 'function') { + let bytes = montgomerySrsMsmBatch( + curveName, + pointBytes, + scalarBytes, + Uint32Array.of(n), + 'single' + ); + if (bytes instanceof Uint8Array && bytes.length === 64 && !isZeroBytes(bytes)) { + return bytes; + } + } + + let msmAffineBytes = curve.Sync?.msmAffineBytes; + if (typeof msmAffineBytes !== 'function') return undefined; + + traceMsmStart(curveName, n, undefined); + let start = traceMsmEnabled() ? performance.now() : 0; + let bytes = msmAffineBytes(pointBytes, scalarBytes); + traceMsm(curveName, n, performance.now() - start, bytes); + if (!(bytes instanceof Uint8Array) || bytes.length !== 64 || isZeroBytes(bytes)) { + return undefined; + } + return bytes; +} + +function montgomerySrsMsmBatch(curveName, pointBytes, scalarBytes, sizes, label) { + if (!isMontgomeryBridgeEnabled() || curves === undefined) return undefined; + let curve = curves[curveName]; + if (curve === undefined) return undefined; + let msmAffineBytesBatch = curve.Sync?.msmAffineBytesBatch; + if (typeof msmAffineBytesBatch !== 'function') return undefined; + + let start = traceMsmEnabled() ? performance.now() : 0; + let bytes; + try { + traceMsmBatchStart(curveName, sizes, label); + bytes = msmAffineBytesBatch(pointBytes, scalarBytes, sizes); + } catch (error) { + traceMsmBatch(curveName, sizes, performance.now() - start, undefined, error, label); + return undefined; + } + traceMsmBatch(curveName, sizes, performance.now() - start, bytes, undefined, label); + if (!(bytes instanceof Uint8Array) || bytes.length !== sizes.length * 64) { + return undefined; + } + return bytes; +} + +function isZeroBytes(bytes) { + for (let i = 0; i < bytes.length; i++) { + if (bytes[i] !== 0) return false; + } + return true; +} + +function isMontgomeryBridgeReady() { + return isMontgomeryBridgeEnabled() && curves !== undefined; +} + +function isMontgomeryProverBridgeReady() { + return isMontgomeryProverMsmEnabled() && isMontgomeryProverPhase() && curves !== undefined; +} + +function isMontgomeryCommitBridgeReady() { + return isMontgomeryCommitMsmEnabled() && curves !== undefined; +} + +function isMontgomeryProverBatchBridgeReady() { + return isMontgomeryProverBatchMsmEnabled() && isMontgomeryProverPhase() && curves !== undefined; +} + +function isMontgomeryBridgeEnabled() { + let globalFlag = globalThis.O1JS_EXPERIMENTAL_MONTGOMERY_MSM; + let processFlag = + typeof process !== 'undefined' ? process.env.O1JS_EXPERIMENTAL_MONTGOMERY_MSM : undefined; + return globalFlag === '1' || globalFlag === true || processFlag === '1'; +} + +function isMontgomeryProverMsmEnabled() { + let globalFlag = globalThis.O1JS_EXPERIMENTAL_MONTGOMERY_PROVER_MSM; + let processFlag = + typeof process !== 'undefined' + ? process.env.O1JS_EXPERIMENTAL_MONTGOMERY_PROVER_MSM + : undefined; + return ( + isMontgomeryBridgeEnabled() && + (globalFlag === '1' || globalFlag === true || processFlag === '1') + ); +} + +function isMontgomeryCommitMsmEnabled() { + let globalFlag = globalThis.O1JS_EXPERIMENTAL_MONTGOMERY_COMMIT_MSM; + let processFlag = + typeof process !== 'undefined' + ? process.env.O1JS_EXPERIMENTAL_MONTGOMERY_COMMIT_MSM + : undefined; + return ( + isMontgomeryBridgeEnabled() && + (globalFlag === '1' || globalFlag === true || processFlag === '1') + ); +} + +function isMontgomeryProverBatchMsmEnabled() { + let globalFlag = globalThis.O1JS_EXPERIMENTAL_MONTGOMERY_PROVER_BATCH_MSM; + let processFlag = + typeof process !== 'undefined' + ? process.env.O1JS_EXPERIMENTAL_MONTGOMERY_PROVER_BATCH_MSM + : undefined; + return ( + isMontgomeryBridgeEnabled() && + (globalFlag === '1' || globalFlag === true || processFlag === '1') + ); +} + +function isMontgomeryProverPhase() { + return globalThis.O1JS_MONTGOMERY_MSM_PHASE === 'prove'; +} + +function traceMsmEnabled() { + return typeof process !== 'undefined' && process.env.O1JS_MONTGOMERY_MSM_TRACE === '1'; +} + +function traceMsm(curve, n, ms, bytes) { + if (!traceMsmEnabled()) return; + let ok = bytes instanceof Uint8Array && bytes.length === 64 && !isZeroBytes(bytes); + process.stderr.write( + `[o1js-montgomery-msm] ${JSON.stringify({ + curve, + phase: currentTracePhase(), + n, + ms, + ok, + thread: typeof process !== 'undefined' ? process.pid : undefined, + })}\n` + ); +} + +function traceMsmStart(curve, n, label) { + if (!traceMsmEnabled()) return; + process.stderr.write( + `[o1js-montgomery-msm-start] ${JSON.stringify({ + curve, + phase: currentTracePhase(), + label, + n, + thread: typeof process !== 'undefined' ? process.pid : undefined, + })}\n` + ); +} + +function traceMsmBatch(curve, sizes, ms, bytes, error, label) { + if (!traceMsmEnabled()) return; + let batchSizes = Array.from(sizes); + let ok = bytes instanceof Uint8Array && bytes.length === batchSizes.length * 64; + process.stderr.write( + `[o1js-montgomery-msm] ${JSON.stringify({ + curve, + phase: currentTracePhase(), + label, + batch: batchSizes.length, + n: batchSizes.reduce((sum, size) => sum + size, 0), + sizes: batchSizes, + ms, + ok, + error: error === undefined ? undefined : String(error?.stack ?? error), + thread: typeof process !== 'undefined' ? process.pid : undefined, + })}\n` + ); +} + +function traceMsmBatchStart(curve, sizes, label) { + if (!traceMsmEnabled()) return; + let batchSizes = Array.from(sizes); + process.stderr.write( + `[o1js-montgomery-msm-start] ${JSON.stringify({ + curve, + phase: currentTracePhase(), + label, + batch: batchSizes.length, + n: batchSizes.reduce((sum, size) => sum + size, 0), + sizes: batchSizes, + thread: typeof process !== 'undefined' ? process.pid : undefined, + })}\n` + ); +} + +function currentTracePhase() { + return globalThis.O1JS_MONTGOMERY_MSM_TRACE_PHASE; +} diff --git a/src/bindings/js/node/node-backend.js b/src/bindings/js/node/node-backend.js index d70477b19d..85d2ab8a20 100644 --- a/src/bindings/js/node/node-backend.js +++ b/src/bindings/js/node/node-backend.js @@ -4,10 +4,15 @@ import { dirname, join } from 'path'; import { fileURLToPath } from 'url'; import { Worker, isMainThread, parentPort, workerData } from 'worker_threads'; import { WithThreadPool, workers } from '../../../lib/proof-system/workers.js'; +import { + installMontgomeryMsmBridge, + isMontgomeryProverMsmEnabled, +} from '../montgomery-msm-bridge.js'; let url = import.meta.url; let filename = url !== undefined ? fileURLToPath(url) : __filename; const require = createRequire(filename); const wasm_ = requireKimchiWasm(!isMainThread ? workerData?.memory : undefined); +const montgomeryBridgeReady = installMontgomeryMsmBridge(); /** * @type {import("../../compiled/node_bindings/kimchi_wasm.cjs")} @@ -18,7 +23,7 @@ if (typeof globalThis !== 'undefined') { globalThis.__o1js_backend_preference = 'wasm'; } -export { wasm, withThreadPool }; +export { montgomeryBridgeReady, wasm, withThreadPool }; function requireKimchiWasm(memoryOverride) { let modulePath = filename.endsWith('index.cjs') @@ -54,6 +59,11 @@ globalThis.startWorkers = startWorkers; globalThis.terminateWorkers = terminateWorkers; if (!isMainThread) { + startWasmWorker(); +} + +async function startWasmWorker() { + if (isMontgomeryProverMsmEnabled()) await montgomeryBridgeReady; parentPort.postMessage({ type: 'wasm_bindgen_worker_ready' }); wasm.wbg_rayon_start_worker(workerData.receiver); } diff --git a/src/bindings/js/web/web-backend.js b/src/bindings/js/web/web-backend.js index 7754edc031..35ef7e6412 100644 --- a/src/bindings/js/web/web-backend.js +++ b/src/bindings/js/web/web-backend.js @@ -1,6 +1,7 @@ import o1jsWebSrc from 'string:../../../web_bindings/o1js_web.bc.js'; import { WithThreadPool, workers } from '../../../lib/proof-system/workers.js'; import kimchiWasm from '../../../web_bindings/kimchi_wasm.js'; +import { installMontgomeryMsmBridge } from '../montgomery-msm-bridge.js'; import { inlineWorker, srcFromFunctionModule, waitForMessage } from './worker-helpers.js'; import { workerSpec } from './worker-spec.js'; @@ -25,6 +26,7 @@ async function initializeBindings() { const memory = allocateWasmMemoryForUserAgent(navigator.userAgent); await init(undefined, memory); + await installMontgomeryMsmBridge(); let module = init.__wbindgen_wasm_module; diff --git a/src/examples/zkprogram/program.ts b/src/examples/zkprogram/program.ts index 7fe3861cb7..67521ce9bb 100644 --- a/src/examples/zkprogram/program.ts +++ b/src/examples/zkprogram/program.ts @@ -1,5 +1,5 @@ -import { Field, ZkProgram, Cache, verify } from 'o1js'; - +import { Cache, Field, Provable, ZkProgram, setBackend, verify } from 'o1js'; +setBackend('wasm'); let MyProgram = ZkProgram({ name: 'example-with-output', publicOutput: Field, @@ -7,6 +7,10 @@ let MyProgram = ZkProgram({ baseCase: { privateInputs: [], async method() { + let a = Provable.witness(Field, () => Field(42)); + for (let index = 0; index < 2 ** 16; index++) { + a = a.mul(a); + } return { publicOutput: Field(1), }; @@ -16,9 +20,9 @@ let MyProgram = ZkProgram({ }); console.time('compile (without cache)'); -let { verificationKey } = await MyProgram.compile({ cache: Cache.None }); +let { verificationKey } = await MyProgram.compile({ cache: Cache.None, forceRecompile: true }); console.timeEnd('compile (without cache)'); - +console.log('constraints', (await MyProgram.analyzeMethods()).baseCase.rows); console.time('proving'); let result = await MyProgram.baseCase(); console.timeEnd('proving'); diff --git a/src/mina b/src/mina index 9e4b7b5fc8..420373d526 160000 --- a/src/mina +++ b/src/mina @@ -1 +1 @@ -Subproject commit 9e4b7b5fc86fcc439286b7da2e5bde57a5114452 +Subproject commit 420373d5268b737f7a1019049cf6cf7ff2f85006