feat: 실시간 음성대화 GPU 엔진(whisper+XTTS) + Ollama 쿼터 자동 폴백
- voice_engine.py: faster-whisper(STT)+XTTS-v2(TTS)를 상시 로드해 로컬에서 서빙하는 WebSocket 엔진 (배치/스트리밍 프로토콜) - routes-voice-realtime.ts: 브라우저 WS를 voice_engine.py로 그대로 프록시 - voice-call.js/worklets: 실시간 연속 대화 모드(VAD 기반 발화 감지, barge-in), 마크다운/표/URL을 정리하고 읽는 sanitizeForSpeech 포함 - tts.ts/stt.ts: xtts_gpu/whisper_gpu provider 분기 추가 - ollama-client.ts/factory.ts: Ollama Cloud 세션 쿼터 초과 시 설정된 fallback 모델로 자동 재시도 Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -62,6 +62,13 @@ export async function transcribeAudio(
|
||||
sourceFormat: string = 'ogg',
|
||||
languageOverride?: string
|
||||
): Promise<{ text: string; language?: string }> {
|
||||
const cfg = getConfig().getConfig() as any;
|
||||
const provider = cfg?.voice?.stt?.provider;
|
||||
if (provider === 'whisper_gpu') {
|
||||
const { transcribeFullGPU } = await import('./voice-engine-client.js');
|
||||
return transcribeFullGPU(audioBuffer, sourceFormat, languageOverride || cfg?.voice?.stt?.language);
|
||||
}
|
||||
|
||||
const config = getSTTConfig();
|
||||
|
||||
if (!config.modelPath) {
|
||||
|
||||
+31
-1
@@ -31,10 +31,40 @@ export function isTTSAvailable(): boolean {
|
||||
}
|
||||
|
||||
export async function synthesizeSpeech(text: string): Promise<Buffer> {
|
||||
const cfg = getConfig().getConfig() as any;
|
||||
const provider = cfg?.voice?.tts?.provider;
|
||||
const truncatedText = text.slice(0, 4000);
|
||||
|
||||
if (provider === 'xtts_gpu') {
|
||||
const tempDir = cfg?.voice?.tempDir || os.tmpdir();
|
||||
const ffmpegPath = cfg?.voice?.ffmpegPath || 'ffmpeg';
|
||||
const { synthesizeFullGPU } = await import('./voice-engine-client.js');
|
||||
const { pcm, sampleRate } = await synthesizeFullGPU(truncatedText);
|
||||
const pcmPath = path.join(tempDir, `tts_gpu_${Date.now()}.pcm`);
|
||||
const oggPath = path.join(tempDir, `tts_gpu_${Date.now()}.ogg`);
|
||||
fs.writeFileSync(pcmPath, pcm);
|
||||
try {
|
||||
await new Promise<void>((resolve, reject) => {
|
||||
const ffmpeg = spawn(ffmpegPath, [
|
||||
'-y', '-f', 's16le', '-ar', String(sampleRate), '-ac', '1', '-i', pcmPath,
|
||||
'-c:a', 'libopus', '-b:a', '48k', '-vbr', 'on', '-compression_level', '10',
|
||||
oggPath,
|
||||
]);
|
||||
let stderr = '';
|
||||
ffmpeg.stderr.on('data', (d: Buffer) => { stderr += d.toString('utf8'); });
|
||||
ffmpeg.on('close', (code) => code === 0 ? resolve() : reject(new Error(`ffmpeg exited ${code}: ${stderr}`)));
|
||||
ffmpeg.on('error', reject);
|
||||
});
|
||||
return fs.readFileSync(oggPath);
|
||||
} finally {
|
||||
try { fs.unlinkSync(pcmPath); } catch {}
|
||||
try { fs.unlinkSync(oggPath); } catch {}
|
||||
}
|
||||
}
|
||||
|
||||
const config = getTTSConfig();
|
||||
if (!config) throw new Error('TTS not configured');
|
||||
|
||||
const truncatedText = text.slice(0, 4000);
|
||||
const tempDir = config.tempDir;
|
||||
const wavPath = path.join(tempDir, `tts_output_${Date.now()}.wav`);
|
||||
const oggPath = path.join(tempDir, `tts_output_${Date.now()}.ogg`);
|
||||
|
||||
@@ -0,0 +1,198 @@
|
||||
import { spawn, ChildProcess } from 'child_process';
|
||||
import path from 'path';
|
||||
import fs from 'fs';
|
||||
import { WebSocket } from 'ws';
|
||||
import { getConfig } from '../config/config.js';
|
||||
|
||||
interface EngineConfig {
|
||||
port: number;
|
||||
sttModel: string;
|
||||
speakerWav: string;
|
||||
device: string;
|
||||
}
|
||||
|
||||
function getEngineConfig(): EngineConfig {
|
||||
const cfg = getConfig().getConfig() as any;
|
||||
const engine = cfg?.voice?.engine || {};
|
||||
return {
|
||||
port: engine.port || 8765,
|
||||
sttModel: engine.sttModel || 'medium',
|
||||
speakerWav: engine.speakerWav || path.join(REPO_ROOT, '.smallclaw', 'voice', 'speaker_ko.wav'),
|
||||
device: engine.device || 'cuda',
|
||||
};
|
||||
}
|
||||
|
||||
const REPO_ROOT = path.join(__dirname, '..', '..');
|
||||
const VENV_PYTHON = path.join(REPO_ROOT, '.smallclaw', 'voice-venv', 'bin', 'python');
|
||||
const ENGINE_SCRIPT = path.join(REPO_ROOT, 'src', 'tools', 'voice_engine.py');
|
||||
|
||||
const STARTUP_TIMEOUT_MS = 180_000; // first load of whisper+XTTS onto GPU can take a while
|
||||
const PING_INTERVAL_MS = 2_000;
|
||||
|
||||
let proc: ChildProcess | null = null;
|
||||
let socket: WebSocket | null = null;
|
||||
let startPromise: Promise<void> | null = null;
|
||||
let restartAttempts = 0;
|
||||
|
||||
type Pending = { resolve: (v: any) => void; reject: (e: any) => void };
|
||||
const pending = new Map<string, Pending>();
|
||||
let nextId = 1;
|
||||
|
||||
function log(...args: any[]) {
|
||||
console.log('[voice-engine-client]', ...args);
|
||||
}
|
||||
|
||||
function resetConnectionState(reason: string) {
|
||||
for (const [, p] of pending) p.reject(new Error(`voice engine connection lost: ${reason}`));
|
||||
pending.clear();
|
||||
socket = null;
|
||||
}
|
||||
|
||||
function spawnEngine(): ChildProcess {
|
||||
const cfg = getEngineConfig();
|
||||
if (!fs.existsSync(VENV_PYTHON)) {
|
||||
throw new Error(`Voice engine venv not found at ${VENV_PYTHON}. Run: python3 -m venv .smallclaw/voice-venv && .smallclaw/voice-venv/bin/pip install faster-whisper coqui-tts g2pkk websockets`);
|
||||
}
|
||||
const args = [
|
||||
ENGINE_SCRIPT,
|
||||
'--port', String(cfg.port),
|
||||
'--stt-model', cfg.sttModel,
|
||||
'--speaker-wav', cfg.speakerWav,
|
||||
'--device', cfg.device,
|
||||
];
|
||||
log('spawning', VENV_PYTHON, args.join(' '));
|
||||
const venvSitePackages = path.join(REPO_ROOT, '.smallclaw', 'voice-venv', 'lib', 'python3.14', 'site-packages');
|
||||
// ctranslate2's GPU wheel dynamically links CUDA 12 libs that aren't present system-wide
|
||||
// (this host runs CUDA 13 + torch's bundled cu130 runtime) — point it at the pip-installed
|
||||
// nvidia-cublas-cu12 / nvidia-cudnn-cu12 packages instead.
|
||||
const cudaLibDirs = [
|
||||
path.join(venvSitePackages, 'nvidia', 'cublas', 'lib'),
|
||||
path.join(venvSitePackages, 'nvidia', 'cudnn', 'lib'),
|
||||
].filter(fs.existsSync);
|
||||
const env = {
|
||||
...process.env,
|
||||
LD_LIBRARY_PATH: [...cudaLibDirs, process.env.LD_LIBRARY_PATH].filter(Boolean).join(':'),
|
||||
COQUI_TOS_AGREED: '1',
|
||||
};
|
||||
const child = spawn(VENV_PYTHON, args, { stdio: ['ignore', 'pipe', 'pipe'], env });
|
||||
child.stdout?.on('data', (d: Buffer) => log('[stdout]', d.toString('utf8').trim()));
|
||||
child.stderr?.on('data', (d: Buffer) => log('[stderr]', d.toString('utf8').trim()));
|
||||
child.on('exit', (code, signal) => {
|
||||
log(`engine process exited code=${code} signal=${signal}`);
|
||||
proc = null;
|
||||
resetConnectionState('process exited');
|
||||
scheduleRestart();
|
||||
});
|
||||
return child;
|
||||
}
|
||||
|
||||
function scheduleRestart() {
|
||||
restartAttempts++;
|
||||
const delay = Math.min(30_000, 2_000 * restartAttempts);
|
||||
log(`scheduling restart in ${delay}ms (attempt ${restartAttempts})`);
|
||||
setTimeout(() => {
|
||||
ensureEngineRunning().catch(e => log('restart failed:', e.message));
|
||||
}, delay);
|
||||
}
|
||||
|
||||
function connectSocket(port: number): Promise<WebSocket> {
|
||||
return new Promise((resolve, reject) => {
|
||||
const ws = new WebSocket(`ws://127.0.0.1:${port}`);
|
||||
const timer = setTimeout(() => { ws.terminate(); reject(new Error('WS connect timeout')); }, STARTUP_TIMEOUT_MS);
|
||||
ws.once('open', () => { clearTimeout(timer); resolve(ws); });
|
||||
ws.once('error', (e) => { clearTimeout(timer); reject(e); });
|
||||
});
|
||||
}
|
||||
|
||||
function wireSocket(ws: WebSocket) {
|
||||
ws.on('message', (data: Buffer, isBinary: boolean) => {
|
||||
if (isBinary) return; // phase 2: streaming binary audio frames
|
||||
let msg: any;
|
||||
try { msg = JSON.parse(data.toString('utf8')); } catch { return; }
|
||||
const p = msg.id != null ? pending.get(String(msg.id)) : undefined;
|
||||
if (!p) return;
|
||||
pending.delete(String(msg.id));
|
||||
if (msg.type === 'error') p.reject(new Error(msg.message || 'voice engine error'));
|
||||
else p.resolve(msg);
|
||||
});
|
||||
ws.on('close', () => resetConnectionState('socket closed'));
|
||||
ws.on('error', (e) => log('socket error:', e.message));
|
||||
}
|
||||
|
||||
export async function ensureEngineRunning(): Promise<void> {
|
||||
if (socket && socket.readyState === WebSocket.OPEN) return;
|
||||
if (startPromise) return startPromise;
|
||||
|
||||
startPromise = (async () => {
|
||||
const cfg = getEngineConfig();
|
||||
if (!proc) proc = spawnEngine();
|
||||
|
||||
let lastErr: any;
|
||||
const deadline = Date.now() + STARTUP_TIMEOUT_MS;
|
||||
while (Date.now() < deadline) {
|
||||
try {
|
||||
const ws = await connectSocket(cfg.port);
|
||||
wireSocket(ws);
|
||||
socket = ws;
|
||||
restartAttempts = 0;
|
||||
log('connected to voice engine on port', cfg.port);
|
||||
return;
|
||||
} catch (e) {
|
||||
lastErr = e;
|
||||
await new Promise(r => setTimeout(r, 1_000));
|
||||
}
|
||||
}
|
||||
throw new Error(`voice engine failed to become ready: ${lastErr?.message || lastErr}`);
|
||||
})();
|
||||
|
||||
try {
|
||||
await startPromise;
|
||||
} finally {
|
||||
startPromise = null;
|
||||
}
|
||||
}
|
||||
|
||||
async function request(msg: Record<string, any>, timeoutMs = 60_000): Promise<any> {
|
||||
await ensureEngineRunning();
|
||||
if (!socket) throw new Error('voice engine socket not connected');
|
||||
const id = String(nextId++);
|
||||
const payload = JSON.stringify({ ...msg, id });
|
||||
return new Promise((resolve, reject) => {
|
||||
const timer = setTimeout(() => {
|
||||
pending.delete(id);
|
||||
reject(new Error('voice engine request timed out'));
|
||||
}, timeoutMs);
|
||||
pending.set(id, {
|
||||
resolve: (v) => { clearTimeout(timer); resolve(v); },
|
||||
reject: (e) => { clearTimeout(timer); reject(e); },
|
||||
});
|
||||
socket!.send(payload);
|
||||
});
|
||||
}
|
||||
|
||||
export async function transcribeFullGPU(audioBuffer: Buffer, format: string, language?: string): Promise<{ text: string; language?: string }> {
|
||||
const res = await request({
|
||||
type: 'stt_transcribe_full',
|
||||
audioBase64: audioBuffer.toString('base64'),
|
||||
format,
|
||||
language: language || 'ko',
|
||||
}, 120_000);
|
||||
return { text: res.text, language: res.language };
|
||||
}
|
||||
|
||||
export async function synthesizeFullGPU(text: string): Promise<{ pcm: Buffer; sampleRate: number }> {
|
||||
const res = await request({ type: 'tts_synthesize_full', text }, 120_000);
|
||||
return { pcm: Buffer.from(res.audioBase64, 'base64'), sampleRate: res.sampleRate };
|
||||
}
|
||||
|
||||
export function isVoiceEngineConfigured(): boolean {
|
||||
return fs.existsSync(VENV_PYTHON) && fs.existsSync(ENGINE_SCRIPT);
|
||||
}
|
||||
|
||||
export function getEnginePort(): number {
|
||||
return getEngineConfig().port;
|
||||
}
|
||||
|
||||
export function isEngineReady(): boolean {
|
||||
return socket !== null && socket.readyState === WebSocket.OPEN;
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Persistent GPU voice engine: loads faster-whisper (STT) and XTTS-v2 (TTS)
|
||||
once and serves them over a localhost WebSocket so Node never pays model
|
||||
load latency per request.
|
||||
|
||||
Batch protocol (used by Telegram + the existing single-shot web TTS button):
|
||||
-> {"type":"ping"}
|
||||
<- {"type":"pong","ready":true}
|
||||
|
||||
-> {"type":"stt_transcribe_full","audioBase64":"...","format":"webm","language":"ko"}
|
||||
<- {"type":"stt_final","text":"...","language":"ko"}
|
||||
|
||||
-> {"type":"tts_synthesize_full","text":"..."}
|
||||
<- {"type":"tts_result","audioBase64":"...","sampleRate":24000}
|
||||
|
||||
Streaming protocol (used by the web UI continuous-conversation mode):
|
||||
-> {"type":"stt_start","sampleRate":16000,"language":"ko"}
|
||||
-> <binary PCM16 mono frames at sampleRate, repeated>
|
||||
<- {"type":"stt_partial","text":"..."} (emitted opportunistically as audio accumulates)
|
||||
-> {"type":"stt_stop"}
|
||||
<- {"type":"stt_final","text":"..."}
|
||||
|
||||
-> {"type":"tts_start","text":"...","id":"..."}
|
||||
<- {"type":"tts_stream_start","sampleRate":24000,"id":"..."}
|
||||
<- <binary PCM16 mono frames, in order>
|
||||
<- {"type":"tts_end","id":"..."}
|
||||
-> {"type":"tts_cancel"} (aborts the in-flight stream started by tts_start)
|
||||
|
||||
<- {"type":"error","message":"..."} on failure (echoes "id" when the request had one)
|
||||
|
||||
One WS connection == one call session; STT and TTS state live in ConnectionState
|
||||
(handle_connection is invoked fresh per client by websockets.serve).
|
||||
"""
|
||||
import argparse
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
import threading
|
||||
|
||||
import numpy as np
|
||||
import websockets
|
||||
|
||||
logging.basicConfig(level=logging.INFO, format='[voice_engine] %(asctime)s %(levelname)s %(message)s')
|
||||
log = logging.getLogger('voice_engine')
|
||||
# faster-whisper logs "Processing audio with duration ..." / "VAD filter
|
||||
# removed ..." at INFO level on every partial decode during streaming STT —
|
||||
# multiple times per second during a live conversation. Keep our own INFO
|
||||
# logs but quiet this one down to WARNING.
|
||||
logging.getLogger('faster_whisper').setLevel(logging.WARNING)
|
||||
|
||||
MIN_PARTIAL_BYTES = int(16000 * 2 * 0.6) # ~0.6s of 16kHz mono PCM16 before we bother decoding
|
||||
|
||||
|
||||
def is_unrecoverable_cuda_error(e: Exception) -> bool:
|
||||
# A CUDA "device-side assert" (seen from XTTS inference_stream on very short
|
||||
# text) poisons the whole CUDA context for the rest of the process — every
|
||||
# GPU call after it fails identically, whisper included. There is no
|
||||
# in-process recovery; the only fix is a fresh CUDA context, i.e. a restart.
|
||||
msg = str(e)
|
||||
return 'CUDA' in msg and ('assert' in msg.lower() or 'device-side' in msg.lower() or 'failed' in msg.lower())
|
||||
|
||||
|
||||
def crash_and_restart(where: str, e: Exception):
|
||||
log.error('Unrecoverable CUDA error in %s: %s — exiting so the supervisor restarts with a fresh CUDA context.', where, e)
|
||||
os._exit(1)
|
||||
|
||||
|
||||
def load_models(args):
|
||||
log.info('Loading faster-whisper model=%s device=%s ...', args.stt_model, args.device)
|
||||
from faster_whisper import WhisperModel
|
||||
stt_model = WhisperModel(args.stt_model, device=args.device, compute_type='float16' if args.device == 'cuda' else 'int8')
|
||||
log.info('faster-whisper ready.')
|
||||
|
||||
log.info('Loading XTTS-v2 ...')
|
||||
from TTS.api import TTS
|
||||
tts_wrapper = TTS('tts_models/multilingual/multi-dataset/xtts_v2').to(args.device)
|
||||
xtts = tts_wrapper.synthesizer.tts_model # underlying TTS.tts.models.xtts.Xtts — has inference_stream()
|
||||
log.info('Computing speaker conditioning latents from %s ...', args.speaker_wav)
|
||||
gpt_cond_latent, speaker_embedding = xtts.get_conditioning_latents(audio_path=[args.speaker_wav])
|
||||
log.info('XTTS-v2 ready.')
|
||||
|
||||
return stt_model, tts_wrapper, xtts, gpt_cond_latent, speaker_embedding
|
||||
|
||||
|
||||
def pcm16_from_float(audio: np.ndarray) -> bytes:
|
||||
clipped = np.clip(audio, -1.0, 1.0)
|
||||
return (clipped * 32767.0).astype('<i2').tobytes()
|
||||
|
||||
|
||||
def pcm16_bytes_to_float(pcm: bytes) -> np.ndarray:
|
||||
return np.frombuffer(pcm, dtype='<i2').astype(np.float32) / 32768.0
|
||||
|
||||
|
||||
class Engine:
|
||||
def __init__(self, args):
|
||||
self.args = args
|
||||
self.stt_model = None
|
||||
self.tts_model = None # TTS.api.TTS wrapper (used for the batch .tts() call)
|
||||
self.xtts = None # underlying Xtts model (used for inference_stream)
|
||||
self.gpt_cond_latent = None
|
||||
self.speaker_embedding = None
|
||||
|
||||
def ready(self) -> bool:
|
||||
return self.stt_model is not None and self.xtts is not None
|
||||
|
||||
# ── Batch STT/TTS (Telegram, single-shot web button) ──────────────────
|
||||
def transcribe_full(self, audio_bytes: bytes, fmt: str, language: str):
|
||||
suffix = '.' + (fmt or 'webm').lstrip('.')
|
||||
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as f:
|
||||
f.write(audio_bytes)
|
||||
path = f.name
|
||||
try:
|
||||
segments, info = self.stt_model.transcribe(path, language=language or None, beam_size=5, vad_filter=True)
|
||||
text = ''.join(seg.text for seg in segments).strip()
|
||||
return text, (info.language if info else language)
|
||||
finally:
|
||||
try:
|
||||
os.unlink(path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
def synthesize_full(self, text: str):
|
||||
wav = self.tts_model.tts(
|
||||
text=text[:4000],
|
||||
speaker_wav=self.args.speaker_wav,
|
||||
language='ko',
|
||||
)
|
||||
audio = np.asarray(wav, dtype=np.float32)
|
||||
pcm = pcm16_from_float(audio)
|
||||
sample_rate = int(self.tts_model.synthesizer.output_sample_rate)
|
||||
return pcm, sample_rate
|
||||
|
||||
# ── Streaming STT (partial decode of an accumulating PCM buffer) ──────
|
||||
def transcribe_pcm(self, pcm: bytes, language: str):
|
||||
audio = pcm16_bytes_to_float(pcm)
|
||||
segments, info = self.stt_model.transcribe(audio, language=language or None, beam_size=1, vad_filter=True)
|
||||
text = ''.join(seg.text for seg in segments).strip()
|
||||
return text, (info.language if info else language)
|
||||
|
||||
# ── Streaming TTS (XTTS inference_stream, chunk-by-chunk) ─────────────
|
||||
def synthesize_stream(self, text: str, is_cancelled):
|
||||
for chunk in self.xtts.inference_stream(
|
||||
text=text[:4000],
|
||||
language='ko',
|
||||
gpt_cond_latent=self.gpt_cond_latent,
|
||||
speaker_embedding=self.speaker_embedding,
|
||||
stream_chunk_size=20,
|
||||
):
|
||||
if is_cancelled():
|
||||
return
|
||||
audio = chunk.squeeze().detach().cpu().numpy().astype(np.float32)
|
||||
yield pcm16_from_float(audio)
|
||||
|
||||
@property
|
||||
def tts_sample_rate(self) -> int:
|
||||
return int(self.tts_model.synthesizer.output_sample_rate)
|
||||
|
||||
|
||||
class ConnectionState:
|
||||
def __init__(self):
|
||||
self.stt_buffer = bytearray()
|
||||
self.stt_decoding = False
|
||||
self.stt_language = 'ko'
|
||||
self.tts_cancel_flags = {} # id -> bool, checked by the producer thread
|
||||
|
||||
|
||||
async def handle_stt_partial(ws, engine: Engine, state: ConnectionState):
|
||||
if state.stt_decoding or len(state.stt_buffer) < MIN_PARTIAL_BYTES:
|
||||
return
|
||||
state.stt_decoding = True
|
||||
try:
|
||||
buf = bytes(state.stt_buffer)
|
||||
loop = asyncio.get_event_loop()
|
||||
text, _lang = await loop.run_in_executor(None, engine.transcribe_pcm, buf, state.stt_language)
|
||||
if text:
|
||||
await ws.send(json.dumps({'type': 'stt_partial', 'text': text}))
|
||||
except Exception as e:
|
||||
log.exception('stt_partial failed')
|
||||
if is_unrecoverable_cuda_error(e):
|
||||
crash_and_restart('handle_stt_partial', e)
|
||||
finally:
|
||||
state.stt_decoding = False
|
||||
|
||||
|
||||
async def handle_tts_start(ws, engine: Engine, state: ConnectionState, msg: dict):
|
||||
text = msg.get('text', '')
|
||||
req_id = msg.get('id')
|
||||
state.tts_cancel_flags[req_id] = False
|
||||
loop = asyncio.get_event_loop()
|
||||
queue: asyncio.Queue = asyncio.Queue()
|
||||
|
||||
def is_cancelled():
|
||||
return state.tts_cancel_flags.get(req_id, False)
|
||||
|
||||
def produce():
|
||||
try:
|
||||
for chunk in engine.synthesize_stream(text, is_cancelled):
|
||||
loop.call_soon_threadsafe(queue.put_nowait, ('chunk', chunk))
|
||||
loop.call_soon_threadsafe(queue.put_nowait, ('done', None))
|
||||
except Exception as e: # noqa: BLE001
|
||||
log.exception('tts stream failed')
|
||||
loop.call_soon_threadsafe(queue.put_nowait, ('error', str(e)))
|
||||
if is_unrecoverable_cuda_error(e):
|
||||
crash_and_restart('handle_tts_start', e)
|
||||
|
||||
threading.Thread(target=produce, daemon=True).start()
|
||||
await ws.send(json.dumps({'type': 'tts_stream_start', 'sampleRate': engine.tts_sample_rate, 'id': req_id}))
|
||||
try:
|
||||
while True:
|
||||
kind, payload = await queue.get()
|
||||
if kind == 'chunk':
|
||||
await ws.send(payload)
|
||||
elif kind == 'error':
|
||||
await ws.send(json.dumps({'type': 'error', 'message': payload, 'id': req_id}))
|
||||
break
|
||||
else:
|
||||
await ws.send(json.dumps({'type': 'tts_end', 'id': req_id}))
|
||||
break
|
||||
finally:
|
||||
state.tts_cancel_flags.pop(req_id, None)
|
||||
|
||||
|
||||
async def handle_connection(ws, engine: Engine):
|
||||
state = ConnectionState()
|
||||
async for raw in ws:
|
||||
try:
|
||||
if isinstance(raw, (bytes, bytearray)):
|
||||
state.stt_buffer.extend(raw)
|
||||
asyncio.create_task(handle_stt_partial(ws, engine, state))
|
||||
continue
|
||||
|
||||
msg = json.loads(raw)
|
||||
mtype = msg.get('type')
|
||||
req_id = msg.get('id')
|
||||
|
||||
if mtype == 'ping':
|
||||
await ws.send(json.dumps({'type': 'pong', 'ready': engine.ready(), 'id': req_id}))
|
||||
|
||||
elif mtype == 'stt_transcribe_full':
|
||||
audio_bytes = base64.b64decode(msg['audioBase64'])
|
||||
fmt = msg.get('format', 'webm')
|
||||
language = msg.get('language', 'ko')
|
||||
loop = asyncio.get_event_loop()
|
||||
text, lang = await loop.run_in_executor(None, engine.transcribe_full, audio_bytes, fmt, language)
|
||||
await ws.send(json.dumps({'type': 'stt_final', 'text': text, 'language': lang, 'id': req_id}))
|
||||
|
||||
elif mtype == 'tts_synthesize_full':
|
||||
text = msg.get('text', '')
|
||||
loop = asyncio.get_event_loop()
|
||||
pcm, sample_rate = await loop.run_in_executor(None, engine.synthesize_full, text)
|
||||
await ws.send(json.dumps({
|
||||
'type': 'tts_result',
|
||||
'audioBase64': base64.b64encode(pcm).decode('ascii'),
|
||||
'sampleRate': sample_rate,
|
||||
'id': req_id,
|
||||
}))
|
||||
|
||||
elif mtype == 'stt_start':
|
||||
state.stt_buffer = bytearray()
|
||||
state.stt_language = msg.get('language', 'ko')
|
||||
|
||||
elif mtype == 'stt_stop':
|
||||
buf = bytes(state.stt_buffer)
|
||||
state.stt_buffer = bytearray()
|
||||
loop = asyncio.get_event_loop()
|
||||
text, lang = ('', state.stt_language)
|
||||
if len(buf) >= 320: # at least ~10ms, skip empty stop
|
||||
text, lang = await loop.run_in_executor(None, engine.transcribe_pcm, buf, state.stt_language)
|
||||
await ws.send(json.dumps({'type': 'stt_final', 'text': text, 'language': lang, 'id': req_id}))
|
||||
|
||||
elif mtype == 'tts_start':
|
||||
asyncio.create_task(handle_tts_start(ws, engine, state, msg))
|
||||
|
||||
elif mtype == 'tts_cancel':
|
||||
target_id = msg.get('id')
|
||||
if target_id is not None:
|
||||
state.tts_cancel_flags[target_id] = True
|
||||
else:
|
||||
for k in state.tts_cancel_flags:
|
||||
state.tts_cancel_flags[k] = True
|
||||
|
||||
else:
|
||||
await ws.send(json.dumps({'type': 'error', 'message': f'unknown type {mtype}', 'id': req_id}))
|
||||
|
||||
except Exception as e:
|
||||
log.exception('request failed')
|
||||
try:
|
||||
await ws.send(json.dumps({'type': 'error', 'message': str(e), 'id': msg.get('id') if 'msg' in dir() else None}))
|
||||
except Exception:
|
||||
pass
|
||||
if is_unrecoverable_cuda_error(e):
|
||||
crash_and_restart('handle_connection', e)
|
||||
|
||||
|
||||
async def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument('--port', type=int, default=8765)
|
||||
parser.add_argument('--stt-model', default='medium')
|
||||
parser.add_argument('--speaker-wav', default='')
|
||||
parser.add_argument('--device', default='cuda')
|
||||
args = parser.parse_args()
|
||||
|
||||
engine = Engine(args)
|
||||
(engine.stt_model, engine.tts_model, engine.xtts,
|
||||
engine.gpt_cond_latent, engine.speaker_embedding) = load_models(args)
|
||||
|
||||
async with websockets.serve(lambda ws: handle_connection(ws, engine), '127.0.0.1', args.port, max_size=64 * 1024 * 1024):
|
||||
log.info('voice_engine listening on 127.0.0.1:%d', args.port)
|
||||
await asyncio.Future()
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
asyncio.run(main())
|
||||
Reference in New Issue
Block a user