feat(omnivoice): streaming /tts + language-safe sanitizer

Add a live-consumer streaming path and text sanitation to the OmniVoice
wrapper, so it can front speech-to-speech chat engines (not just the
asset-engine's batch WAV use).

- POST /tts: chunked 24 kHz mono s16le PCM (or open-ended WAV), driven by
  the adaptive buffer-ratchet scheduler. Emits the first sentence
  immediately, then ratchets chunk size up on OmniVoice's ~40x realtime
  headroom -> sub-second time-to-first-audio. Wire-compatible with
  chatterbox-fast /tts (both 24 kHz mono PCM). Batch /v1/audio/speech is
  unchanged for asset/file callers.

- scheduler.py: VENDORED byte-faithful copy of chatterbox-fast's pure-
  Python (torch-free) scheduler, pinned to commit 7631462 (v0.1.0/v0.1.1).
  Vendor-copy over a shared package (operator call 2026-06-19): the module
  has no GPU deps, so reuse it without dragging chatterbox-fast's torch
  tree into this image. Promote to a shared package only on a 3rd consumer
  or real drift.

- sanitize.py: language-safe TTS sanitizer run on both endpoints. Strips
  markdown, <think> blocks, HTML, and model control tokens; deliberately
  SKIPS the fork's English-only number/phone normalization that would
  corrupt OmniVoice's 600-language input. Preserves [laughter]-style tags.

- Refactor: shared GenParams base for SpeechRequest + TTSStreamRequest;
  single GEN_LOCK serializes generation (single-stream interactive).

- Dockerfile/playbook: copy + upload the two new modules; build-time
  `import app` smoke; correct stale "Gradio demo / no FastAPI" comments.
This commit is contained in:
vh
2026-06-19 22:47:15 -07:00
parent 826c2a6a64
commit 288d085236
7 changed files with 671 additions and 67 deletions
+219 -55
View File
@@ -1,5 +1,4 @@
"""Thin FastAPI wrapper exposing OmniVoice (k2-fsa/OmniVoice) as an OpenAI-style
TTS for the asset-engine.
"""Thin FastAPI wrapper exposing OmniVoice (k2-fsa/OmniVoice) for the fleet.
Upstream ships only a Gradio demo; we own this wrapper (same pattern as
stacks/index-tts/app.py). Voices are reference WAVs staged in
@@ -7,35 +6,55 @@ ${OMNIVOICE_VOICES_DIR} (the reused chatterbox /refs/*.wav). A voice-clone promp
is precomputed once per voice at startup (the loaded Whisper ASR auto-transcribes
each reference) and cached, so per-request latency is just generation.
Two consumption modes:
- BATCH (asset-engine / OpenAI-compat): POST /v1/audio/speech -> one WAV blob.
- STREAM (live speech-to-speech chat engines): POST /tts -> chunked PCM, driven
by the vendored adaptive buffer-ratchet scheduler (scheduler.py, from
chatterbox-fast). Emits the first sentence immediately for sub-second
time-to-first-audio, then ratchets chunk size up on OmniVoice's ~40x realtime
headroom. Wire-compatible with chatterbox-fast's /tts (both 24 kHz mono s16le).
All text is run through the language-safe sanitizer (sanitize.py) before synthesis
on BOTH endpoints — strips markdown / LLM artifacts / control tokens without the
English-only normalization that would corrupt OmniVoice's multilingual input.
Exposes OmniVoice's full generation surface:
- clone (voice=<staged ref>) and/or voice-DESIGN (instruct=<free-text style>)
- clone (voice=<staged ref>) and/or voice-DESIGN (instruct=<controlled tags>)
- language (Auto + 600+), speed, duration
- diffusion controls: num_step, guidance_scale, denoise, preprocess_prompt,
postprocess_output, plus a generation_overrides passthrough for expert knobs
(t_shift, layer_penalty_factor, position_temperature, class_temperature, ...).
Endpoints:
GET /healthz -> readiness (200 once model + >=1 voice are loaded)
GET /v1/audio/voices -> {"voices": [<name>, ...]}
GET /v1/audio/languages -> {"languages": ["Auto", <display name>, ...]}
POST /v1/audio/speech -> audio/wav
GET /healthz -> readiness (200 once model + >=1 voice are loaded)
GET /v1/audio/voices -> {"voices": [<name>, ...]}
GET /v1/audio/languages -> {"languages": ["Auto", <display name>, ...]}
GET /v1/audio/instruct-items -> {"instruct_items": [<tag>, ...]}
POST /v1/audio/speech -> audio/wav (batch, OpenAI-style)
POST /tts -> streaming PCM/WAV (chatterbox-fast-compatible)
"""
import glob
import io
import logging
import os
import struct
import threading
import time
from pathlib import Path
from typing import Any, Dict, Optional
from typing import Any, Dict, Literal, Optional
import numpy as np
import soundfile as sf
import torch
from fastapi import FastAPI, HTTPException
from fastapi.responses import Response
from pydantic import BaseModel
from fastapi.responses import Response, StreamingResponse
from pydantic import BaseModel, Field
from omnivoice import OmniVoice, OmniVoiceGenerationConfig
from sanitize import sanitize_tts_text
from scheduler import ChunkConfig, ChunkResult, stream_chunks
try:
from omnivoice.utils.lang_map import LANG_NAMES, lang_display_name
LANGUAGES = ["Auto"] + sorted(lang_display_name(n) for n in LANG_NAMES)
@@ -57,18 +76,25 @@ CKPT = os.environ.get("OMNIVOICE_CKPT", "k2-fsa/OmniVoice")
VOICES_DIR = os.environ.get("OMNIVOICE_VOICES_DIR", "/app/voices")
ASR_MODEL = os.environ.get("OMNIVOICE_ASR_MODEL", "openai/whisper-large-v3-turbo")
app = FastAPI(title="OmniVoice TTS (asset-engine wrapper)")
app = FastAPI(title="OmniVoice TTS (asset-engine + streaming wrapper)")
MODEL: Optional[OmniVoice] = None
PROMPTS: dict = {} # voice name -> VoiceClonePrompt
SR: int = 24000
# Generation is serialized: the workload is single-stream interactive and a
# streaming request holds the model for the duration of its stream. Concurrent
# callers queue rather than interleave on the GPU.
GEN_LOCK = threading.Lock()
class GenParams(BaseModel):
"""OmniVoice generation parameters shared by the batch and streaming endpoints."""
class SpeechRequest(BaseModel):
input: str
# Voice source — at least one of voice (clone) / instruct (design) is required.
voice: Optional[str] = None # staged reference clip -> clone timbre
instruct: Optional[str] = None # free-text voice DESIGN / style
instruct: Optional[str] = None # voice DESIGN / style (controlled tags)
# Generation controls (defaults mirror the upstream demo).
language: Optional[str] = "Auto" # "Auto" -> auto-detect
speed: Optional[float] = None # 0.5–1.5; ignored if duration set
@@ -82,10 +108,27 @@ class SpeechRequest(BaseModel):
# layer_penalty_factor, position_temperature, class_temperature,
# audio_chunk_duration, audio_chunk_threshold). Unknown keys are dropped.
generation_overrides: Optional[Dict[str, Any]] = None
class SpeechRequest(GenParams):
"""OpenAI-style /v1/audio/speech (batch) request."""
response_format: str = "wav"
model: Optional[str] = None # ignored (single model); OpenAI-compat
class TTSStreamRequest(GenParams):
"""Streaming /tts request — chatterbox-fast-compatible wire protocol."""
format: Literal["pcm", "wav"] = "pcm" # raw s16le PCM (default) or open-ended WAV
stream: bool = True # False -> whole-text one-shot (A/B vs stream)
# Scheduler overrides (None -> ChunkConfig defaults; see scheduler.py).
margin: Optional[float] = Field(default=None)
margin_first: Optional[float] = Field(default=None)
rtf_prior: Optional[float] = Field(default=None)
sec_per_char_prior: Optional[float] = Field(default=None)
@app.on_event("startup")
def _load() -> None:
global MODEL, SR
@@ -108,6 +151,104 @@ def _load() -> None:
log.info("%d voices loaded; %d languages", len(PROMPTS), len(LANGUAGES))
# ── generation helpers (shared by batch + streaming) ────────────────────────
def _base_gen_kwargs(req: GenParams) -> Dict[str, Any]:
"""Validate the voice source and build the MODEL.generate kwargs minus `text`.
Raises HTTPException (400/404) for caller errors — call this BEFORE a stream
starts so those land as proper status codes, not mid-stream failures.
"""
has_instruct = bool(req.instruct and req.instruct.strip())
if not req.voice and not has_instruct:
raise HTTPException(
status_code=400,
detail="provide a 'voice' (clone a staged reference) and/or 'instruct' (design a voice)",
)
cfg = {
"num_step": req.num_step,
"guidance_scale": req.guidance_scale,
"denoise": req.denoise,
"preprocess_prompt": req.preprocess_prompt,
"postprocess_output": req.postprocess_output,
**(req.generation_overrides or {}),
}
kw: Dict[str, Any] = {"generation_config": OmniVoiceGenerationConfig.from_dict(cfg)}
if req.voice:
prompt = PROMPTS.get(req.voice)
if prompt is None:
raise HTTPException(
status_code=404,
detail=f"unknown voice '{req.voice}'; have {sorted(PROMPTS)}",
)
kw["voice_clone_prompt"] = prompt
if has_instruct:
kw["instruct"] = req.instruct.strip()
if req.language and req.language != "Auto":
kw["language"] = req.language
if req.speed is not None:
kw["speed"] = req.speed
if req.duration is not None:
kw["duration"] = req.duration
return kw
def _synth(text: str, base_kw: Dict[str, Any]) -> "tuple[np.ndarray, float]":
"""Synthesize one text span -> (float32 audio [-1,1], audio_seconds)."""
out = MODEL.generate(text=text, **base_kw)
audio = out[0] if isinstance(out, (list, tuple)) else out
audio = np.asarray(audio, dtype=np.float32)
return audio, (len(audio) / SR if SR else 0.0)
def _pcm16(audio: np.ndarray) -> bytes:
"""float32 [-1,1] -> little-endian s16 PCM bytes (24 kHz mono on the wire)."""
a = np.clip(np.asarray(audio, dtype=np.float32).reshape(-1), -1.0, 1.0)
return (a * 32767.0).astype("<i2").tobytes()
def _wav_header(sr: int, data_len: Optional[int] = None) -> bytes:
"""WAV header. data_len=None -> streaming (0xFFFFFFFF sizes, read to EOF);
an int -> correct RIFF/data sizes for a complete file."""
data_size = 0xFFFFFFFF if data_len is None else data_len
riff_size = 0xFFFFFFFF if data_len is None else 36 + data_len
return (
b"RIFF" + struct.pack("<I", riff_size) + b"WAVE"
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, sr, sr * 2, 2, 16)
+ b"data" + struct.pack("<I", data_size)
)
def _chunk_config(req: TTSStreamRequest) -> ChunkConfig:
cfg = ChunkConfig()
if req.margin is not None:
cfg.margin = req.margin
if req.margin_first is not None:
cfg.margin_first = req.margin_first
if req.rtf_prior is not None:
cfg.rtf_prior = req.rtf_prior
if req.sec_per_char_prior is not None:
cfg.sec_per_char_prior = req.sec_per_char_prior
return cfg
def _log_chunk(r: ChunkResult, ttfa_ms: float) -> None:
tag = " STARVED" if r.starved else ""
if r.index == 0:
log.info("chunk 0: ttfa=%.0fms gen=%.0fms audio=%.2fs rtf=%.2f%s",
ttfa_ms, r.gen_time * 1000, r.audio_sec, r.rtf, tag)
else:
log.info("chunk %d: gen=%.0fms audio=%.2fs buf=%.2f->%.2f rtf=%.2f%s",
r.index, r.gen_time * 1000, r.audio_sec,
r.buffer_before, r.buffer_after, r.rtf, tag)
# ── read-only discovery endpoints ───────────────────────────────────────────
@app.get("/healthz")
def healthz():
if MODEL is None or not PROMPTS:
@@ -132,59 +273,82 @@ def instruct_items():
return {"instruct_items": INSTRUCT_ITEMS}
# ── synthesis endpoints ─────────────────────────────────────────────────────
@app.post("/v1/audio/speech")
def speech(req: SpeechRequest):
"""Batch (OpenAI-style): generate the whole utterance, return one WAV blob."""
if MODEL is None:
raise HTTPException(status_code=503, detail="model not loaded")
if not req.input or not req.input.strip():
raise HTTPException(status_code=400, detail="input is empty")
if req.response_format not in ("wav", "", None):
raise HTTPException(status_code=400, detail="only response_format=wav is supported")
text = sanitize_tts_text(req.input)
if not text:
raise HTTPException(status_code=400, detail="input is empty after sanitization")
has_instruct = bool(req.instruct and req.instruct.strip())
if not req.voice and not has_instruct:
raise HTTPException(
status_code=400,
detail="provide a 'voice' (clone a staged reference) and/or 'instruct' (design a voice)",
)
cfg = {
"num_step": req.num_step,
"guidance_scale": req.guidance_scale,
"denoise": req.denoise,
"preprocess_prompt": req.preprocess_prompt,
"postprocess_output": req.postprocess_output,
**(req.generation_overrides or {}),
}
gen = OmniVoiceGenerationConfig.from_dict(cfg)
kw: Dict[str, Any] = {"text": req.input.strip(), "generation_config": gen}
if req.voice:
prompt = PROMPTS.get(req.voice)
if prompt is None:
raise HTTPException(
status_code=404,
detail=f"unknown voice '{req.voice}'; have {sorted(PROMPTS)}",
)
kw["voice_clone_prompt"] = prompt
if has_instruct:
kw["instruct"] = req.instruct.strip()
if req.language and req.language != "Auto":
kw["language"] = req.language
if req.speed is not None:
kw["speed"] = req.speed
if req.duration is not None:
kw["duration"] = req.duration
base_kw = _base_gen_kwargs(req) # 400/404 propagate as-is
try:
out = MODEL.generate(**kw)
with GEN_LOCK:
audio, _ = _synth(text, base_kw)
except HTTPException:
raise
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=400, detail=f"{type(exc).__name__}: {exc}")
audio = out[0] if isinstance(out, (list, tuple)) else out
audio = np.asarray(audio, dtype=np.float32)
buf = io.BytesIO()
sf.write(buf, audio, SR, format="WAV", subtype="PCM_16")
return Response(content=buf.getvalue(), media_type="audio/wav")
@app.post("/tts")
def tts(req: TTSStreamRequest):
"""Streaming: chunked 24 kHz mono PCM (or open-ended WAV) for live consumers.
Wire-compatible with chatterbox-fast /tts. `stream=true` (default) runs the
buffer-ratchet schedule for sub-second time-to-first-audio; `stream=false`
is a whole-text one-shot for A/B comparison.
"""
if MODEL is None:
raise HTTPException(status_code=503, detail="model not loaded")
text = sanitize_tts_text(req.input)
if not text:
raise HTTPException(status_code=400, detail="input is empty after sanitization")
base_kw = _base_gen_kwargs(req) # validate up front (pre-stream)
cfg = _chunk_config(req)
media = "audio/wav" if req.format == "wav" else "application/octet-stream"
def body():
# One request holds the lock for its whole stream (single-stream workload).
with GEN_LOCK:
t_req = time.perf_counter()
ttfa_ms: Optional[float] = None
if not req.stream:
audio, audio_sec = _synth(text, base_kw)
ttfa_ms = (time.perf_counter() - t_req) * 1000
log.info("oneshot: %.0fms gen, %.2fs audio", ttfa_ms, audio_sec)
pcm = _pcm16(audio)
if req.format == "wav":
yield _wav_header(SR, len(pcm))
yield pcm
return
if req.format == "wav":
yield _wav_header(SR) # open-ended; length unknown up front
def _gen(t: str):
return _synth(t, base_kw)
total_audio = 0.0
for r in stream_chunks(text, generate=_gen, clock=time.perf_counter, cfg=cfg):
if ttfa_ms is None:
ttfa_ms = (time.perf_counter() - t_req) * 1000
total_audio += r.audio_sec
_log_chunk(r, ttfa_ms)
yield _pcm16(r.audio)
log.info("stream done: ttfa=%.0fms total_audio=%.2fs",
ttfa_ms or 0.0, total_audio)
return StreamingResponse(body(), media_type=media)