Luigit
repositories / dotfiles

dotfiles

bugabingas dorkfiles

owned by admin

quickshell/nuguland/noko-parakeet.mjs

Raw
// noko-parakeet.mjs — Parakeet ONNX model-directory runtime.

import { createRequire } from 'node:module';
import { existsSync, readFileSync, readdirSync } from 'node:fs';
import { join } from 'node:path';
import { loadOnnxRuntime } from './noko-vad.mjs';

const require = createRequire(import.meta.url);
const DEFAULT_MAX_TOKENS_PER_STEP = 10;

export function loadSentencePiece() {
  try {
    return require('sentencepiece-js');
  } catch (error) {
    throw new Error(`tokenizer runtime unavailable: sentencepiece-js (${error.message})`);
  }
}

function walk(dir, prefix = '') {
  if (!dir || !existsSync(dir)) return [];
  const out = [];
  for (const entry of readdirSync(dir, { withFileTypes: true })) {
    if (entry.name.startsWith('._')) continue;
    const rel = prefix ? `${prefix}/${entry.name}` : entry.name;
    const abs = join(dir, entry.name);
    if (entry.isDirectory()) out.push(...walk(abs, rel));
    else out.push({ rel, abs });
  }
  return out;
}

function pick(files, tests) {
  return files.find(file => tests.some(test => test.test(file.rel.toLowerCase()))) || null;
}

export function discoverParakeetDir(dir) {
  const files = walk(dir);
  const found = {
    preprocessor: pick(files, [/preprocessor.*\.onnx$/, /nemo\d+.*\.onnx$/]),
    encoder: pick(files, [/encoder.*\.onnx$/]),
    decoderJoint: pick(files, [/(decoder|joint).*\.onnx$/]),
    tokenizer: pick(files, [/tokenizer\.model$/, /vocab.*\.(txt|json)$/])
  };
  const missing = [];
  if (!found.preprocessor) missing.push('preprocessor');
  if (!found.encoder) missing.push('encoder');
  if (!found.decoderJoint) missing.push('decoder-or-joint');
  if (!found.tokenizer) missing.push('vocabulary-or-tokenizer');
  return {
    path: dir || '',
    ok: missing.length === 0,
    missing,
    files: Object.fromEntries(Object.entries(found).map(([key, value]) => [key, value ? value.rel : '']))
  };
}

export async function loadTokenizer(path, runtime) {
  const sp = runtime || loadSentencePiece();
  const processor = new sp.SentencePieceProcessor();
  await processor.load(path);
  return processor;
}

export function decodeTokens(tokenizer, tokens) {
  if (!tokenizer || typeof tokenizer.decodeIds !== 'function') throw new Error('sentencepiece tokenizer missing');
  return tokenizer.decodeIds((tokens || []).map(Number));
}

function loadVocabulary(path) {
  const vocab = new Map();
  for (const line of readFileSync(path, 'utf8').split(/\r?\n/)) {
    if (!line) continue;
    const split = line.lastIndexOf(' ');
    if (split <= 0) continue;
    vocab.set(Number(line.slice(split + 1)), line.slice(0, split).replaceAll('▁', ' '));
  }
  const blank = Array.from(vocab.entries()).find(([, token]) => token === '<blk>');
  return {
    type: 'vocab',
    vocab,
    size: vocab.size,
    blankIdx: blank ? blank[0] : Math.max(...vocab.keys()),
    decodeIds(ids) {
      const text = ids.map(id => vocab.get(id) || '').join('');
      return text.replace(/^\s|\s\B|(\s)\b/g, (_match, keep) => keep ? ' ' : '');
    }
  };
}

async function loadDecoder(path, options) {
  if (path.toLowerCase().endsWith('.model')) {
    const tokenizer = await loadTokenizer(path, options.sentencepiece);
    return {
      type: 'sentencepiece',
      size: null,
      blankIdx: null,
      decodeIds: ids => decodeTokens(tokenizer, ids)
    };
  }
  return loadVocabulary(path);
}

function first(record, names) {
  for (const name of names) if (record[name] !== undefined) return record[name];
  return Object.values(record)[0];
}

function tensorData(tensor) {
  return tensor && tensor.data !== undefined ? tensor.data : tensor;
}

function numericDim(value, fallback) {
  return typeof value === 'number' && value > 0 ? value : fallback;
}

function argmax(data, start = 0, end = data.length) {
  let best = start;
  for (let i = start + 1; i < end; i++) if (data[i] > data[best]) best = i;
  return best - start;
}

export class ParakeetOnnx {
  constructor(ort, sessions, model, decoder, options = {}) {
    this.ort = ort;
    this.sessions = sessions;
    this.model = model;
    this.decoder = decoder;
    this.maxTokensPerStep = Number(options.maxTokensPerStep || DEFAULT_MAX_TOKENS_PER_STEP);
    this.busy = false;
    this.state = 'ready';
  }

  static async create(dir, options = {}) {
    const model = discoverParakeetDir(dir);
    if (!model.ok) throw new Error(`parakeet model dir incomplete: ${model.missing.join(', ')}`);
    const ort = options.ort || await loadOnnxRuntime();
    const sessionOptions = { executionProviders: options.executionProviders || ['cpu'] };
    const sessions = {
      preprocessor: await ort.InferenceSession.create(join(dir, model.files.preprocessor), sessionOptions),
      encoder: await ort.InferenceSession.create(join(dir, model.files.encoder), sessionOptions),
      decoderJoint: await ort.InferenceSession.create(join(dir, model.files.decoderJoint), sessionOptions)
    };
    const decoder = await loadDecoder(join(dir, model.files.tokenizer), options);
    return new ParakeetOnnx(ort, sessions, model, decoder, options);
  }

  async preprocess(samples) {
    const output = await this.sessions.preprocessor.run({
      waveforms: new this.ort.Tensor('float32', samples, [1, samples.length]),
      waveforms_lens: new this.ort.Tensor('int64', BigInt64Array.from([BigInt(samples.length)]), [1])
    });
    return {
      features: first(output, ['features']),
      lengths: first(output, ['features_lens'])
    };
  }

  async encode(features, lengths) {
    const output = await this.sessions.encoder.run({
      audio_signal: features,
      length: lengths
    });
    return {
      encoded: first(output, ['outputs']),
      lengths: first(output, ['encoded_lengths'])
    };
  }

  createDecoderState() {
    const inputs = Object.fromEntries((this.sessions.decoderJoint.inputNames || []).map(name => [name, null]));
    const metadata = Object.fromEntries((this.sessions.decoderJoint.inputMetadata || []).map(input => [input.name, input]));
    const shape1 = metadata.input_states_1 ? metadata.input_states_1.shape : [2, 1, 640];
    const shape2 = metadata.input_states_2 ? metadata.input_states_2.shape : [2, 1, 640];
    void inputs;
    return {
      state1: new Float32Array(numericDim(shape1[0], 2) * 1 * numericDim(shape1[2], 640)),
      state2: new Float32Array(numericDim(shape2[0], 2) * 1 * numericDim(shape2[2], 640)),
      state1Shape: [numericDim(shape1[0], 2), 1, numericDim(shape1[2], 640)],
      state2Shape: [numericDim(shape2[0], 2), 1, numericDim(shape2[2], 640)]
    };
  }

  async decodeFrame(prevTokens, state, frame) {
    const previous = prevTokens.length ? prevTokens[prevTokens.length - 1] : this.decoder.blankIdx;
    const output = await this.sessions.decoderJoint.run({
      encoder_outputs: new this.ort.Tensor('float32', frame, [1, frame.length, 1]),
      targets: new this.ort.Tensor('int32', Int32Array.from([previous]), [1, 1]),
      target_length: new this.ort.Tensor('int32', Int32Array.from([1]), [1]),
      input_states_1: new this.ort.Tensor('float32', state.state1, state.state1Shape),
      input_states_2: new this.ort.Tensor('float32', state.state2, state.state2Shape)
    });
    const logits = tensorData(first(output, ['outputs']));
    return {
      logits,
      step: argmax(logits, this.decoder.size, logits.length),
      state: {
        state1: new Float32Array(tensorData(first(output, ['output_states_1']))),
        state2: new Float32Array(tensorData(first(output, ['output_states_2']))),
        state1Shape: state.state1Shape,
        state2Shape: state.state2Shape
      }
    };
  }

  async transcribe(samples) {
    if (this.busy) throw new Error('parakeet inference already running');
    if (!samples || samples.length === 0) throw new Error('empty speech buffer');
    this.busy = true;
    this.state = 'transcribing';
    try {
      const { features, lengths } = await this.preprocess(samples instanceof Float32Array ? samples : new Float32Array(samples));
      const { encoded, lengths: encodedLengths } = await this.encode(features, lengths);
      const encodedData = tensorData(encoded);
      const dims = encoded.dims || encoded.dimensions;
      const channels = Number(dims[1]);
      const frames = Number(dims[2]);
      const limit = Math.min(frames, Number(tensorData(encodedLengths)[0]));
      const state = this.createDecoderState();
      let currentState = state;
      const tokens = [];
      let emittedTokens = 0;
      for (let t = 0; t < limit;) {
        const frame = new Float32Array(channels);
        for (let c = 0; c < channels; c++) frame[c] = encodedData[c * frames + t];
        const decoded = await this.decodeFrame(tokens, currentState, frame);
        const token = argmax(decoded.logits, 0, this.decoder.size);
        if (token !== this.decoder.blankIdx) {
          currentState = decoded.state;
          tokens.push(token);
          emittedTokens++;
        }
        if (decoded.step > 0) {
          t += decoded.step;
          emittedTokens = 0;
        } else if (token === this.decoder.blankIdx || emittedTokens === this.maxTokensPerStep) {
          t++;
          emittedTokens = 0;
        }
      }
      return {
        text: this.decoder.decodeIds(tokens),
        meta: {
          model: 'parakeet-onnx',
          quantization: this.model.files.encoder.includes('int8') ? 'int8' : 'unknown',
          runtimeProvider: 'onnxruntime-node'
        }
      };
    } finally {
      this.busy = false;
      this.state = 'ready';
    }
  }
}