cd92b85157
Empirical follow-up to the streaming /tts smoke test on the 3090. OmniVoice is diffusion: a ~fixed per-call overhead (~1.5s at 32 steps, ~0.7s at 16) dominates regardless of chunk length, so the upstream-claimed 40x RTF does NOT hold here (measured ~2.8x/32-step, ~5.6x/16-step) and the chatterbox- tuned scheduler over-chunks and starves. - Streaming /tts defaults to num_step=16 (TTFA ~1.5s -> ~0.7s); batch /v1/audio/speech stays num_step=32 for quality. Per-request override intact. - Scheduler prior raised to rtf_prior=20 (env OMNIVOICE_STREAM_RTF_PRIOR, wired through compose + .env.example) so it packs whole-text-minus-first- sentence into a few chunks: validated ~3 chunks, no starvation, total wall ~= one-shot, less per-chunk silence padding. - Docs corrected: the "sub-second / 40x" claims were wrong; streaming has a diffusion TTFA floor (~0.7s) and wins mainly on long replies. chatterbox- fast (autoregressive, ~0.5s TTFA) stays the lowest-latency front-end; OmniVoice is the multilingual / voice-design complement.
370 lines
15 KiB
Python
370 lines
15 KiB
Python
"""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
|
||
${OMNIVOICE_VOICES_DIR} (the reused chatterbox /refs/*.wav). A voice-clone prompt
|
||
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 so first-audio comes
|
||
sooner than one-shot, then packs the rest into a few chunks. NB: OmniVoice is
|
||
diffusion, so a ~fixed per-call overhead sets a TTFA floor (~0.7s at 16 steps
|
||
on the 3090, NOT sub-second); the win grows with utterance length. 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=<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>, ...]}
|
||
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, Literal, Optional
|
||
|
||
import numpy as np
|
||
import soundfile as sf
|
||
import torch
|
||
from fastapi import FastAPI, HTTPException
|
||
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)
|
||
except Exception: # noqa: BLE001
|
||
LANGUAGES = ["Auto"]
|
||
|
||
try:
|
||
# Voice-DESIGN `instruct` is a CONTROLLED vocabulary (gender / age / pitch /
|
||
# accent / whisper tags), NOT free prose — surfaced so callers can discover it.
|
||
from omnivoice.utils.voice_design import _INSTRUCT_VALID_EN
|
||
INSTRUCT_ITEMS = sorted(_INSTRUCT_VALID_EN)
|
||
except Exception: # noqa: BLE001
|
||
INSTRUCT_ITEMS = []
|
||
|
||
logging.basicConfig(level=os.environ.get("OMNIVOICE_LOG_LEVEL", "INFO"))
|
||
log = logging.getLogger("omnivoice-api")
|
||
|
||
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")
|
||
|
||
# Streaming scheduler prior. OmniVoice is diffusion: a ~fixed per-call overhead
|
||
# dominates (short and long chunks cost ~the same), so the chatterbox default
|
||
# (rtf_prior=3.4) over-chunks and STARVES — each extra chunk re-pays the fixed
|
||
# cost and adds boundary silence. A high prior packs whole-text-minus-first-
|
||
# sentence into a few chunks (validated on the 3090: ~3 chunks, no starvation,
|
||
# total ≈ one-shot). Per-request `rtf_prior` still overrides this.
|
||
OMNIVOICE_STREAM_RTF_PRIOR = float(os.environ.get("OMNIVOICE_STREAM_RTF_PRIOR", "20"))
|
||
|
||
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."""
|
||
|
||
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 # 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
|
||
duration: Optional[float] = None # fixed seconds; overrides speed
|
||
num_step: int = 32 # 4–64 diffusion steps (batch=32; /tts overrides to 16)
|
||
guidance_scale: float = 2.0 # 0.0–4.0 CFG
|
||
denoise: bool = True
|
||
preprocess_prompt: bool = True
|
||
postprocess_output: bool = True
|
||
# Expert passthrough into OmniVoiceGenerationConfig (t_shift,
|
||
# 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."""
|
||
|
||
# Streaming defaults to FEWER diffusion steps than batch (32): halves the
|
||
# ~per-call diffusion overhead (server-side TTFA ~1.5s -> ~0.7s on the 3090)
|
||
# at some quality cost. Override per-request for the quality/latency trade.
|
||
num_step: int = 16
|
||
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
|
||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||
log.info("loading OmniVoice %s on %s (asr=%s)", CKPT, device, ASR_MODEL)
|
||
MODEL = OmniVoice.from_pretrained(
|
||
CKPT, device_map=device, load_asr=True, asr_model_name=ASR_MODEL
|
||
)
|
||
SR = int(getattr(MODEL, "sampling_rate", 24000))
|
||
|
||
for wav in sorted(glob.glob(os.path.join(VOICES_DIR, "*.wav"))):
|
||
name = Path(wav).stem
|
||
if name.startswith("_"):
|
||
continue # skip _*.wav (chatterbox test/deploy artifacts)
|
||
try:
|
||
PROMPTS[name] = MODEL.create_voice_clone_prompt(ref_audio=wav)
|
||
log.info("voice ready: %s", name)
|
||
except Exception as exc: # noqa: BLE001
|
||
log.warning("voice %s failed to load: %s", name, exc)
|
||
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()
|
||
cfg.rtf_prior = OMNIVOICE_STREAM_RTF_PRIOR # diffusion-aware default (pack aggressively)
|
||
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:
|
||
raise HTTPException(status_code=503, detail="not ready")
|
||
return {"status": "ok", "voices": len(PROMPTS), "languages": len(LANGUAGES),
|
||
"sampling_rate": SR}
|
||
|
||
|
||
@app.get("/v1/audio/voices")
|
||
def voices():
|
||
return {"voices": sorted(PROMPTS.keys())}
|
||
|
||
|
||
@app.get("/v1/audio/languages")
|
||
def languages():
|
||
return {"languages": LANGUAGES}
|
||
|
||
|
||
@app.get("/v1/audio/instruct-items")
|
||
def instruct_items():
|
||
# Valid comma-separable voice-DESIGN attribute tags (English).
|
||
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 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")
|
||
|
||
base_kw = _base_gen_kwargs(req) # 400/404 propagate as-is
|
||
try:
|
||
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}")
|
||
|
||
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)
|