码桶

发现社区成员的开源项目

奶狗 /

ng-webot

公开
main
ng-webot/lib/stt.js
stt.js6.1 KB
'use strict';

/**
 * 本地语音转文字(纯本地推理,不依赖任何外部接口)。
 * 使用 @huggingface/transformers 在 Node 端跑 Whisper(ONNX),
 * 模型权重首次自动下载并缓存到 data/stt-models,之后完全离线。
 *
 * 输入要求:单声道、16000Hz、[-1,1] 浮点采样(Float32Array)。
 */

const fs = require('fs');
const path = require('path');
const silk = require('silk-wasm');

const MODELS_DIR = path.join(__dirname, '..', 'data', 'stt-models');
fs.mkdirSync(MODELS_DIR, { recursive: true });

// 模型缓存目录;下载源默认走国内镜像(HuggingFace 官方域名在国内常被墙)。
// 如需换源(如海外服务器),设置环境变量 STT_HF_ENDPOINT=https://huggingface.co/
let _transformers = null;
function getTransformers() {
  if (!_transformers) {
    try {
      _transformers = require('@huggingface/transformers');
    } catch (e) {
      console.error('[stt] 加载 @huggingface/transformers 失败(可选依赖),语音转文字不可用:', e.message);
      throw e;
    }
    _transformers.env.cacheDir = MODELS_DIR;
    _transformers.env.allowRemoteModels = true;
    _transformers.env.remoteHost = process.env.STT_HF_ENDPOINT || 'https://hf-mirror.com/';
  }
  return _transformers;
}

let _transcriberPromise = null;
let _transcriberModel = null;

async function getTranscriber(modelId) {
  if (_transcriberPromise && _transcriberModel === modelId) return _transcriberPromise;
  const { pipeline } = getTransformers();
  _transcriberModel = modelId;
  console.log('[stt] 加载本地语音识别模型:', modelId, '(首次会自动下载,请耐心等待)');
  _transcriberPromise = pipeline('automatic-speech-recognition', modelId, {
    progress_callback: (p) => {
      if (p && p.status === 'progress' && p.total) {
        const pct = Math.round((p.loaded / p.total) * 100);
        console.log(`[stt] 下载 ${p.file || ''} ${pct}%`);
      }
    },
  }).catch(e => {
    console.error('[stt] 模型加载失败,重置缓存以便下次重试:', e.message);
    _transcriberPromise = null;
    _transcriberModel = null;
    throw e;
  });
  return _transcriberPromise;
}

// ---------- 音频解码 / 重采样 ----------

/** 线性插值重采样:input(Float32) 从 inRate 重采样到 outRate */
function resampleLinear(input, inRate, outRate) {
  if (inRate === outRate) return input;
  const ratio = inRate / outRate;
  const newLen = Math.max(1, Math.round(input.length / ratio));
  const out = new Float32Array(newLen);
  for (let i = 0; i < newLen; i++) {
    const pos = i * ratio;
    const i0 = Math.floor(pos);
    const i1 = Math.min(i0 + 1, input.length - 1);
    const frac = pos - i0;
    out[i] = input[i0] * (1 - frac) + input[i1] * frac;
  }
  return out;
}

/** Int16 PCM → Float32([-1,1]) */
function int16ToFloat32(int16) {
  const f = new Float32Array(int16.length);
  for (let i = 0; i < int16.length; i++) f[i] = int16[i] / 32768;
  return f;
}

/** 从 SILK 解码为 Int16 PCM(按 ref.sample_rate 或 24000 解码,再重采样到 16000) */
async function silkToFloat32(buffer, ref) {
  const nativeRate = (ref && ref.sample_rate) || 24000;
  const { data } = await silk.decode(buffer, nativeRate);
  if (!data || !data.byteLength) throw new Error('SILK 解码结果为空');
  const int16 = new Int16Array(data.buffer, data.byteOffset, data.byteLength >> 1);
  const f = int16ToFloat32(int16);
  return resampleLinear(f, nativeRate, 16000);
}

/** 解析 WAV 为 Int16 PCM(仅支持 16bit PCM) */
function parseWavInt16(buffer) {
  if (buffer.readUInt32LE(0) !== 0x46464952 || buffer.readUInt32LE(8) !== 0x45564157) {
    throw new Error('不是有效的 WAV 文件');
  }
  let pos = 12;
  let sampleRate = 16000, bits = 16, channels = 1;
  let dataStart = -1, dataLen = 0;
  while (pos + 8 <= buffer.length) {
    const id = buffer.toString('ascii', pos, pos + 4);
    const size = buffer.readUInt32LE(pos + 4);
    if (id === 'fmt ') {
      const af = buffer.readUInt16LE(pos + 8);
      channels = buffer.readUInt16LE(pos + 10);
      sampleRate = buffer.readUInt32LE(pos + 12);
      bits = buffer.readUInt16LE(pos + 22);
    } else if (id === 'data') {
      dataStart = pos + 8;
      dataLen = size;
      break;
    }
    pos += 8 + size + (size & 1);
  }
  if (dataStart < 0) throw new Error('WAV 缺少 data 块');
  if (bits !== 16) throw new Error('仅支持 16bit PCM WAV');
  let int16 = new Int16Array(buffer.buffer, buffer.byteOffset + dataStart, dataLen >> 1);
  if (channels > 1) {
    const mono = new Int16Array(int16.length / channels);
    for (let i = 0; i < mono.length; i++) {
      let s = 0;
      for (let c = 0; c < channels; c++) s += int16[i * channels + c];
      mono[i] = s / channels;
    }
    int16 = mono;
  }
  return { int16, sampleRate };
}

/** 把微信入站语音(解密后的原始字节)解码为 16000Hz 单声道 Float32 */
async function decodeVoiceToFloat(buffer, ref) {
  if (silk.isWav(buffer)) {
    const { int16, sampleRate } = parseWavInt16(buffer);
    return resampleLinear(int16ToFloat32(int16), sampleRate, 16000);
  }
  if (silk.isSilk(buffer)) {
    return await silkToFloat32(buffer, ref);
  }
  throw new Error('本地识别仅支持微信语音(SILK)或 WAV,其它格式请改用 AI 接口模式');
}

// ---------- 对外接口 ----------

/**
 * 本地 Whisper 转写。
 * @param {Float32Array} audio 16000Hz 单声道浮点采样
 * @param {{model?:string, language?:string}} [opts]
 * @returns {Promise<string>} 识别文本
 */
async function transcribeLocal(audio, opts = {}) {
  const model = opts.model || process.env.STT_MODEL || 'Xenova/whisper-base';
  const language = opts.language || 'chinese';
  const transcriber = await getTranscriber(model);
  const output = await transcriber(audio, {
    language,
    task: 'transcribe',
    return_timestamps: false,
  });
  const text = output && output.text ? output.text : '';
  return String(text).trim();
}

module.exports = {
  MODELS_DIR,
  decodeVoiceToFloat,
  transcribeLocal,
};