#!/usr/bin/env node const fs = require('node:fs'); const path = require('node:path'); const ort = require('onnxruntime-web'); const root = path.resolve(__dirname, '..'); const modelId = process.argv[2]; if (!modelId) throw new Error('usage: node scripts/validate_web.cjs '); ort.env.wasm.numThreads = 1; ort.env.wasm.proxy = false; async function main() { const catalog = JSON.parse(fs.readFileSync(path.join(root, 'catalog.json'), 'utf8')); const cases = JSON.parse(fs.readFileSync(path.join(root, 'validation/cases.json'), 'utf8')); const model = catalog.models.find((item) => item.id === modelId); if (!model) throw new Error(`unknown model: ${modelId}`); const tokenizer = catalog.tokenizers.find((item) => item.id === model.tokenizerId); const config = JSON.parse(fs.readFileSync(path.join(root, tokenizer.artifact.path), 'utf8')); const vocab = config[catalog.runtime.tokenEncoding.vocabularyField]; let language; let voice; for (const candidateLanguage of catalog.languages) { const candidateVoice = candidateLanguage.voices.find((item) => item.modelId === modelId); if (candidateVoice) { language = candidateLanguage; voice = candidateVoice; break; } } if (!voice) throw new Error(`no voice references model: ${modelId}`); const testCase = cases.find((item) => { if (!item.webModelTest) return false; const candidateLanguage = catalog.languages.find((entry) => entry.id === item.language); const candidateVoice = candidateLanguage?.voices.find((entry) => entry.id === item.voice); return candidateVoice?.modelId === modelId; }); if (!testCase) throw new Error(`no webModelTest case for model: ${modelId}`); language = catalog.languages.find((item) => item.id === testCase.language); voice = language.voices.find((item) => item.id === testCase.voice); function idsFor(phonemes) { return [ catalog.runtime.tokenEncoding.bosTokenId, ...Array.from(phonemes).map((phoneme) => { const id = vocab[phoneme]; if (id === undefined) throw new Error(`missing phoneme: ${phoneme}`); return id; }), catalog.runtime.tokenEncoding.eosTokenId, ]; } const voiceBytes = fs.readFileSync(path.join(root, voice.artifact.path)); function styleAt(row) { const styleOffset = row * 256 * 4; const style = new Float32Array(256); for (let index = 0; index < style.length; index += 1) { style[index] = voiceBytes.readFloatLE(styleOffset + index * 4); } return style; } const modelBytes = fs.readFileSync(path.join(root, model.artifact.path)); const session = await ort.InferenceSession.create(modelBytes, { executionProviders: ['wasm'], graphOptimizationLevel: 'all', }); async function validate(caseId, ids, style) { const output = await session.run({ input_ids: new ort.Tensor( 'int64', BigInt64Array.from(ids.map(BigInt)), [1, ids.length], ), style: new ort.Tensor('float32', style, [1, 256]), speed: new ort.Tensor('float32', Float32Array.of(1), [1]), }); const waveform = output.waveform.data; const durationSum = Array.from(output.duration.data).reduce( (sum, item) => sum + Number(item), 0, ); const expectedSamples = durationSum * catalog.runtime.durationToSamplesFactor; if (waveform.length !== expectedSamples) { throw new Error(`waveform/duration mismatch: ${waveform.length} vs ${expectedSamples}`); } let peak = 0; for (const sample of waveform) { if (!Number.isFinite(sample)) throw new Error('non-finite waveform'); peak = Math.max(peak, Math.abs(sample)); } return { id: caseId, inputTokens: ids.length, outputSamples: waveform.length, peakAbsoluteAmplitude: peak, }; } const phonemeCount = Array.from(testCase.phonemes).length; const testIds = idsFor(testCase.phonemes); const maxPhonemes = catalog.runtime.maxPhonemeCodePoints; const maxIds = idsFor('a'.repeat(maxPhonemes)); const result = { modelId, language: language.id, voice: voice.id, onnxRuntimeWebVersion: JSON.parse( fs.readFileSync(path.join(root, 'node_modules/onnxruntime-web/package.json'), 'utf8'), ).version, executionProvider: 'wasm', cases: [ await validate(testCase.id, testIds, styleAt(phonemeCount - 1)), await validate('max-phoneme-boundary', maxIds, styleAt(maxPhonemes - 1)), ], }; const filename = `web-${modelId.replace(/[^a-z0-9_-]/g, '-')}.json`; fs.writeFileSync( path.join(root, 'validation', filename), `${JSON.stringify(result, null, 2)}\n`, ); console.log(JSON.stringify(result)); } main().catch((error) => { console.error(error); process.exitCode = 1; });