UNPKG

hama-js

Version:

G2P, phoneme-ASR, and P2G inference for Node, Bun, and browsers, powered by a self-contained WASM engine (no onnxruntime).

181 lines 7.79 kB
// WASM-backed inference engine (the compiled Zig `hama.wasm`). Marshals inputs // into the module's linear memory, calls the exported model functions, and reads // outputs back. Used identically by the Node/Bun and browser entry points. // // NOTE: the linear memory may grow (and its ArrayBuffer detach) on any alloc or // run call, so typed-array views are always created from the *current* // `memory.buffer` immediately before use. export const ENC_FEAT = 192; export const ENC_HID = 96; export const DEC_HID = 96; export const ASR_VOCAB = 191; export class HamaEngine { constructor(ex) { this.ex = ex; } static async fromBytes(wasm) { const { instance } = await WebAssembly.instantiate(wasm, {}); return new HamaEngine(instance.exports); } writeBytes(ptr, src) { new Uint8Array(this.ex.memory.buffer, ptr, src.length).set(src); } writeF32(ptr, src) { new Float32Array(this.ex.memory.buffer, ptr, src.length).set(src); } writeI64(ptr, src) { new BigInt64Array(this.ex.memory.buffer, ptr, src.length).set(src); } readF32(ptr, len) { return new Float32Array(this.ex.memory.buffer, ptr, len).slice(); } readU8(ptr, len) { return new Uint8Array(this.ex.memory.buffer, ptr, len).slice(); } readI64(ptr, len) { return new BigInt64Array(this.ex.memory.buffer, ptr, len).slice(); } loadModel(kind, bytes) { const ptr = this.ex.hama_alloc(bytes.length); if (ptr === 0) throw new Error("hama_alloc failed"); this.writeBytes(ptr, bytes); const fn = kind === "encoder" ? this.ex.hama_encoder_load : kind === "decoder" ? this.ex.hama_decoder_load : kind === "asr" ? this.ex.hama_asr_load : this.ex.hama_p2g_load; const h = fn(ptr, bytes.length); this.ex.hama_free(ptr, bytes.length); if (h === 0) throw new Error(`hama_${kind}_load failed`); return h; } loadEncoder(bytes) { return this.loadModel("encoder", bytes); } loadDecoder(bytes) { return this.loadModel("decoder", bytes); } loadAsr(bytes) { return this.loadModel("asr", bytes); } loadP2g(bytes) { return this.loadModel("p2g", bytes); } /** Greedy decode: prefixIds = [bos, src, phones..., tgt]; returns generated token ids. */ p2gGreedy(h, prefixIds, maxNew, eos, pad) { const P = prefixIds.length; const pPtr = this.ex.hama_alloc(P * 8); const outPtr = this.ex.hama_alloc(maxNew * 8); this.writeI64(pPtr, prefixIds); const n = Number(this.ex.hama_p2g_greedy(h, pPtr, P, maxNew, BigInt(eos), BigInt(pad), outPtr)); if (n < 0) throw new Error("hama_p2g_greedy failed"); const out = this.readI64(outPtr, maxNew); this.ex.hama_free(pPtr, P * 8); this.ex.hama_free(outPtr, maxNew * 8); return Array.from(out.slice(0, n), Number); } /** Greedy decode + per-token source-phoneme alignment index (-1 if unaligned). */ p2gGreedyAlign(h, prefixIds, maxNew, eos, pad) { const P = prefixIds.length; const pPtr = this.ex.hama_alloc(P * 8); const outPtr = this.ex.hama_alloc(maxNew * 8); const alignPtr = this.ex.hama_alloc(maxNew * 8); this.writeI64(pPtr, prefixIds); const n = Number(this.ex.hama_p2g_greedy_align(h, pPtr, P, maxNew, BigInt(eos), BigInt(pad), outPtr, alignPtr)); if (n < 0) throw new Error("hama_p2g_greedy_align failed"); const ids = Array.from(this.readI64(outPtr, maxNew).slice(0, n), Number); const align = Array.from(this.readI64(alignPtr, maxNew).slice(0, n), Number); this.ex.hama_free(pPtr, P * 8); this.ex.hama_free(outPtr, maxNew * 8); this.ex.hama_free(alignPtr, maxNew * 8); return { ids, align }; } encoderRun(h, ids, length) { const T = ids.length; const idsPtr = this.ex.hama_alloc(T * 8); const eoPtr = this.ex.hama_alloc(T * ENC_FEAT * 4); const pkPtr = this.ex.hama_alloc(T * ENC_HID * 4); const hidPtr = this.ex.hama_alloc(2 * ENC_HID * 4); const maskPtr = this.ex.hama_alloc(T); const prevPtr = this.ex.hama_alloc(T * 4); this.writeI64(idsPtr, ids); const rc = this.ex.hama_encoder_run(h, idsPtr, T, length, eoPtr, pkPtr, hidPtr, maskPtr, prevPtr); if (rc !== 0) throw new Error("hama_encoder_run failed"); const out = { encoderOutputs: this.readF32(eoPtr, T * ENC_FEAT), projectedKeys: this.readF32(pkPtr, T * ENC_HID), hidden: this.readF32(hidPtr, 2 * ENC_HID), mask: this.readU8(maskPtr, T), prevAttn: this.readF32(prevPtr, T), T, }; this.ex.hama_free(idsPtr, T * 8); this.ex.hama_free(eoPtr, T * ENC_FEAT * 4); this.ex.hama_free(pkPtr, T * ENC_HID * 4); this.ex.hama_free(hidPtr, 2 * ENC_HID * 4); this.ex.hama_free(maskPtr, T); this.ex.hama_free(prevPtr, T * 4); return out; } decoderStep(h, token, eo, pk, mask, prev, hidden, positions) { const T = prev.length; const eoPtr = this.ex.hama_alloc(eo.length * 4); const pkPtr = this.ex.hama_alloc(pk.length * 4); const maskPtr = this.ex.hama_alloc(T); const prevPtr = this.ex.hama_alloc(T * 4); const hidPtr = this.ex.hama_alloc(hidden.length * 4); const posPtr = this.ex.hama_alloc(T * 4); const nextPtr = this.ex.hama_alloc(8); const attnPtr = this.ex.hama_alloc(8); const hidOutPtr = this.ex.hama_alloc(2 * DEC_HID * 4); const prevOutPtr = this.ex.hama_alloc(T * 4); this.writeF32(eoPtr, eo); this.writeF32(pkPtr, pk); this.writeBytes(maskPtr, mask); this.writeF32(prevPtr, prev); this.writeF32(hidPtr, hidden); this.writeF32(posPtr, positions); const rc = this.ex.hama_decoder_step(h, BigInt(token), eoPtr, pkPtr, maskPtr, prevPtr, hidPtr, posPtr, T, nextPtr, attnPtr, hidOutPtr, prevOutPtr); if (rc !== 0) throw new Error("hama_decoder_step failed"); const out = { nextToken: Number(this.readI64(nextPtr, 1)[0]), attnArgmax: Number(this.readI64(attnPtr, 1)[0]), hiddenOut: this.readF32(hidOutPtr, 2 * DEC_HID), prevOut: this.readF32(prevOutPtr, T), }; this.ex.hama_free(eoPtr, eo.length * 4); this.ex.hama_free(pkPtr, pk.length * 4); this.ex.hama_free(maskPtr, T); this.ex.hama_free(prevPtr, T * 4); this.ex.hama_free(hidPtr, hidden.length * 4); this.ex.hama_free(posPtr, T * 4); this.ex.hama_free(nextPtr, 8); this.ex.hama_free(attnPtr, 8); this.ex.hama_free(hidOutPtr, 2 * DEC_HID * 4); this.ex.hama_free(prevOutPtr, T * 4); return out; } asrRun(h, wav) { const N = wav.length; const T = this.ex.hama_asr_num_frames(N); const wavPtr = this.ex.hama_alloc(N * 4); const lpPtr = this.ex.hama_alloc(T * ASR_VOCAB * 4); const olPtr = this.ex.hama_alloc(8); this.writeF32(wavPtr, wav); const rc = this.ex.hama_asr_run(h, wavPtr, N, lpPtr, olPtr); if (rc < 0) throw new Error("hama_asr_run failed"); const logProbs = this.readF32(lpPtr, T * ASR_VOCAB); const outLength = Number(this.readI64(olPtr, 1)[0]); this.ex.hama_free(wavPtr, N * 4); this.ex.hama_free(lpPtr, T * ASR_VOCAB * 4); this.ex.hama_free(olPtr, 8); return { logProbs, T, outLength }; } } //# sourceMappingURL=engine.js.map