feat(stt): real Korean STT via persistent faster-whisper worker
Step 3 (귀): add WhisperSTT + whisper_worker, a warm out-of-venv worker mirroring the MeloTTS shape (whisper312 venv, small/int8 on CPU). transcribe() closes the voice round trip (MeloTTS wav -> whisper text); utterances() turns an injected audio_source into Utterances (Discord voice feed pending). Wired into factory as WSAI_STT=whisper. Also address the arbiter's TTS follow-ups: - melo worker error handling: capture stderr (drained in a bounded background task so the pipe can't fill), surface the real failure cause, and defend against an empty/invalid ready line instead of dying on JSONDecodeError. - pipeline pre-warm: load slow backends (warmup()) at startup so the first utterance is answered warm; a warmup failure is logged, not fatal. Verified: real TTS->STT round trip recovers the sentence near-perfectly; warm transcribe ~1.2s (CPU). 12 tests pass. Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
81
wsai/backends/whisper_worker.py
Normal file
81
wsai/backends/whisper_worker.py
Normal file
@@ -0,0 +1,81 @@
|
||||
"""Persistent faster-whisper STT worker.
|
||||
|
||||
faster-whisper (ctranslate2) lives in its own Python (whisper312); loading the
|
||||
model takes seconds, so we load it ONCE here and then serve transcription
|
||||
requests over stdin/stdout. This process is launched with the whisper312
|
||||
interpreter by wsai.backends.whisper.WhisperSTT.
|
||||
|
||||
Like the MeloTTS worker, model/backend chatter could corrupt the JSON protocol,
|
||||
so on startup we split the streams: a private duplicate of the original stdout
|
||||
carries the protocol, and fd 1 is redirected to fd 2 so any library print lands
|
||||
on stderr instead (where the parent drains it for diagnostics).
|
||||
|
||||
Protocol (one JSON object per line, on the protocol channel):
|
||||
<- {"wav": "/abs/path.wav", "language": "ko"}
|
||||
-> {"ok": true, "text": "...", "language": "ko", "ms": 123}
|
||||
-> {"ok": false, "error": "..."}
|
||||
On startup, once the model is ready, it emits exactly one line:
|
||||
-> {"ready": true, "ms": <load-ms>, "device": "cpu", "model": "small"}
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
# Split protocol from library noise BEFORE importing anything heavy.
|
||||
_proto = os.fdopen(os.dup(1), "w", buffering=1) # private copy of real stdout
|
||||
os.dup2(2, 1) # fd1 -> stderr, so stray library prints don't hit the protocol
|
||||
|
||||
|
||||
def _emit(obj: dict) -> None:
|
||||
_proto.write(json.dumps(obj, ensure_ascii=False) + "\n")
|
||||
_proto.flush()
|
||||
|
||||
|
||||
def _log(*a):
|
||||
print(*a, file=sys.stderr, flush=True)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
model_name = os.environ.get("WSAI_WHISPER_MODEL", "small")
|
||||
device = os.environ.get("WSAI_WHISPER_DEVICE", "cpu") # "cpu" | "cuda"
|
||||
# int8 on CPU keeps a small model fast; float16 is the usual CUDA choice.
|
||||
compute = os.environ.get(
|
||||
"WSAI_WHISPER_COMPUTE", "int8" if device == "cpu" else "float16"
|
||||
)
|
||||
default_lang = os.environ.get("WSAI_WHISPER_LANGUAGE", "ko") or None
|
||||
|
||||
t0 = time.monotonic()
|
||||
from faster_whisper import WhisperModel # heavy import; only in whisper venv
|
||||
|
||||
model = WhisperModel(model_name, device=device, compute_type=compute)
|
||||
load_ms = int((time.monotonic() - t0) * 1000)
|
||||
_emit({"ready": True, "ms": load_ms, "device": device, "model": model_name})
|
||||
_log(f"[whisper_worker] {model_name} ready in {load_ms} ms on {device}/{compute}")
|
||||
|
||||
for line in sys.stdin:
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
req = json.loads(line)
|
||||
wav = req["wav"]
|
||||
language = req.get("language", default_lang)
|
||||
s = time.monotonic()
|
||||
segments, info = model.transcribe(
|
||||
wav,
|
||||
language=language,
|
||||
beam_size=int(req.get("beam_size", 5)),
|
||||
vad_filter=bool(req.get("vad_filter", True)),
|
||||
)
|
||||
text = "".join(seg.text for seg in segments).strip()
|
||||
ms = int((time.monotonic() - s) * 1000)
|
||||
_emit({"ok": True, "text": text, "language": info.language, "ms": ms})
|
||||
except Exception as exc: # keep the worker alive across bad requests
|
||||
_emit({"ok": False, "error": f"{type(exc).__name__}: {exc}"})
|
||||
_log(f"[whisper_worker] error: {exc}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user