63 lines
2.6 KiB
JavaScript
63 lines
2.6 KiB
JavaScript
/* Can we get real language detection out of Whisper in Transformers.js? */
|
|
import { AutoProcessor, WhisperForConditionalGeneration, Tensor, env } from '@huggingface/transformers';
|
|
import { synthesize } from '../src/speech/index.js';
|
|
|
|
env.cacheDir = './.transformers-cache';
|
|
const MODEL = 'onnx-community/whisper-base';
|
|
|
|
const processor = await AutoProcessor.from_pretrained(MODEL);
|
|
const model = await WhisperForConditionalGeneration.from_pretrained(MODEL, { dtype: 'q8' });
|
|
const tok = processor.tokenizer;
|
|
|
|
// Whisper emits one language token right after <|startoftranscript|>. Reading
|
|
// that distribution is a single decoder step — far cheaper than transcribing
|
|
// twice to see which language "looks better".
|
|
const id = (t) => tok.encode(t, { add_special_tokens: false })[0];
|
|
const sot = id('<|startoftranscript|>');
|
|
const CANDIDATES = ['en', 'ta'];
|
|
const langIds = CANDIDATES.map((c) => id(`<|${c}|>`));
|
|
console.log('sot:', sot, '| language token ids:', JSON.stringify(Object.fromEntries(CANDIDATES.map((c, i) => [c, langIds[i]]))));
|
|
|
|
function resample(a, from, to) {
|
|
const r = from / to;
|
|
const out = new Float32Array(Math.floor(a.length / r));
|
|
for (let i = 0; i < out.length; i++) {
|
|
const p = i * r, k = Math.floor(p);
|
|
out[i] = a[k] + (a[Math.min(k + 1, a.length - 1)] - a[k]) * (p - k);
|
|
}
|
|
return out;
|
|
}
|
|
|
|
for (const [lang, text] of [
|
|
['ta', 'மூவாயிரம் நானூற்று இருபத்தேழு புதிய லீட்கள் உள்ளன.'],
|
|
['en', 'There are three thousand four hundred and twenty seven new leads.'],
|
|
]) {
|
|
const spoken = await synthesize(text, lang);
|
|
const audio = resample(spoken.audio, spoken.sampling_rate, 16000);
|
|
|
|
const inputs = await processor(audio);
|
|
const t0 = Date.now();
|
|
const out = await model({
|
|
...inputs,
|
|
decoder_input_ids: new Tensor('int64', BigInt64Array.from([BigInt(sot)]), [1, 1]),
|
|
});
|
|
const ms = Date.now() - t0;
|
|
|
|
const logits = out.logits;
|
|
const last = logits.dims[1] - 1;
|
|
const vocab = logits.dims[2];
|
|
const row = logits.data.slice(last * vocab, (last + 1) * vocab);
|
|
|
|
const scores = langIds.map((id) => Number(row[id]));
|
|
const max = Math.max(...scores);
|
|
const exp = scores.map((s) => Math.exp(s - max));
|
|
const sum = exp.reduce((a, b) => a + b, 0);
|
|
const probs = exp.map((e) => e / sum);
|
|
const best = probs.indexOf(Math.max(...probs));
|
|
|
|
console.log(`spoken ${lang} → detected ${CANDIDATES[best]} `
|
|
+ `(${CANDIDATES.map((c, i) => `${c} ${probs[i].toFixed(3)}`).join(', ')}) in ${ms}ms `
|
|
+ `${CANDIDATES[best] === lang ? '✅' : '❌'}`);
|
|
}
|
|
process.exit(0);
|