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:
+219
-55
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user