Luigit
repositories / dotfiles

dotfiles

bugabingas dorkfiles

owned by admin

quickshell/nuguland/noko-vad.mjs

Raw
// noko-vad.mjs — real ONNX VAD runtime bridge for Noko.
// Loads Silero/equivalent ONNX through onnxruntime-node when available.

import { createRequire } from 'node:module';
import { existsSync } from 'node:fs';

const require = createRequire(import.meta.url);
const SILERO_CHUNK_16K = 512;
const SILERO_CONTEXT_16K = 64;
const SILERO_STATE_SIZE = 128;

function dataDir() {
  if (process.env.NOKO_DATA_DIR) return process.env.NOKO_DATA_DIR;
  const root = process.env.XDG_DATA_HOME || `${process.env.HOME || '.'}/.local/share`;
  return `${root}/nuguland/noko`;
}

function runtimeModulePath() {
  return process.env.NOKO_ONNXRUNTIME_NODE || `${dataDir()}/runtime/node_modules/onnxruntime-node`;
}

export async function loadOnnxRuntime() {
  try {
    return require('onnxruntime-node');
  } catch (requireError) {
    try {
      return await import('onnxruntime-node');
    } catch (importError) {
      try {
        return require(runtimeModulePath());
      } catch (localError) {
        throw new Error(`runtime provider unavailable: onnxruntime-node (${requireError.message}; ${importError.message}; ${localError.message})`);
      }
    }
  }
}

function tensorValue(tensor) {
  if (!tensor) return null;
  if (tensor.data !== undefined) return tensor.data;
  return tensor;
}

function scalar(value) {
  const data = tensorValue(value);
  if (ArrayBuffer.isView(data) || Array.isArray(data)) return Number(data[0] || 0);
  return Number(data || 0);
}

function firstPresent(record, names) {
  for (const name of names) if (record && record[name] !== undefined) return record[name];
  return null;
}

export async function vadRuntimeReport(modelPath, options = {}) {
  const report = {
    path: modelPath || '',
    loadable: false,
    inputNames: [],
    outputNames: [],
    stateful: false,
    srInput: false,
    frameSamples: 480,
    smokeProbability: null,
    error: ''
  };
  try {
    const vad = await OnnxVad.create(modelPath, options);
    report.loadable = true;
    report.inputNames = vad.session.inputNames || [];
    report.outputNames = vad.session.outputNames || [];
    report.stateful = report.inputNames.indexOf('state') >= 0;
    report.srInput = report.inputNames.indexOf('sr') >= 0;
    if (options.smoke !== false)
      report.smokeProbability = await vad.probability(new Float32Array(vad.frameSamples));
  } catch (error) {
    report.error = String(error && error.message ? error.message : error);
  }
  return report;
}

export class OnnxVad {
  constructor(ort, session, options = {}) {
    if (!ort || !ort.Tensor) throw new Error('onnxruntime Tensor constructor missing');
    if (!session || typeof session.run !== 'function') throw new Error('onnxruntime session missing');
    this.ort = ort;
    this.session = session;
    this.sampleRate = Number(options.sampleRate || 16000);
    this.frameSamples = Number(options.frameSamples || 480);
    this.state = new Float32Array(2 * SILERO_STATE_SIZE);
    this.context = new Float32Array(SILERO_CONTEXT_16K);
    this.inputNames = session.inputNames || ['input'];
  }

  static async create(modelPath, options = {}) {
    if (!modelPath || !existsSync(modelPath)) throw new Error(`vad model missing: ${modelPath || '(unset)'}`);
    const ort = options.ort || await loadOnnxRuntime();
    const session = options.session || await ort.InferenceSession.create(modelPath, {
      executionProviders: options.executionProviders || ['cpu']
    });
    return new OnnxVad(ort, session, options);
  }

  reset() {
    this.state.fill(0);
    this.context.fill(0);
  }

  sileroInput(frame) {
    const out = new Float32Array(SILERO_CONTEXT_16K + SILERO_CHUNK_16K);
    out.set(this.context, 0);
    out.set(frame.slice(0, this.frameSamples), SILERO_CONTEXT_16K);
    this.context = out.slice(out.length - SILERO_CONTEXT_16K);
    return out;
  }

  feeds(frame) {
    if (!frame || frame.length !== this.frameSamples) throw new Error(`vad frame must be exactly ${this.frameSamples} samples`);
    const hasState = this.inputNames.indexOf('state') >= 0;
    const input = hasState ? this.sileroInput(frame) : new Float32Array(frame);
    const feeds = {
      input: new this.ort.Tensor('float32', input, [1, input.length])
    };
    if (hasState) feeds.state = new this.ort.Tensor('float32', this.state, [2, 1, SILERO_STATE_SIZE]);
    if (this.inputNames.indexOf('sr') >= 0) feeds.sr = new this.ort.Tensor('int64', BigInt64Array.from([BigInt(this.sampleRate)]), []);
    return feeds;
  }

  async probability(frame) {
    const output = await this.session.run(this.feeds(frame));
    const nextState = firstPresent(output, ['stateN', 'state', 'hn', 'h']);
    if (nextState) this.state = new Float32Array(tensorValue(nextState));
    return scalar(firstPresent(output, ['output', 'prob', 'speech_probs']) || Object.values(output)[0]);
  }
}

export async function vadProbabilities(samples, options = {}) {
  const vad = options.vad || await OnnxVad.create(options.modelPath || process.env.NOKO_VAD_ONNX || '', options);
  const frameSamples = options.frameSamples || 480;
  const out = [];
  for (let offset = 0; offset + frameSamples <= samples.length; offset += frameSamples) {
    out.push(await vad.probability(samples.slice(offset, offset + frameSamples)));
  }
  return out;
}