feat(dots-tts): ship OpenAI-compatible dots.tts TTS stack on irv-ml1:8198
Thin FastAPI wrapper over DotsTtsRuntime (soar, optimize=True, RTF ~0.22), serialized single-consumer; OpenAI /v1/audio/speech (stream + non-stream), voices from the voices/ corpus derived set. Live + healthy alongside chatterbox-fast on the 3090; nothing repointed. Dockerfile needs build-essential (torch.compile/inductor JITs via gcc at runtime) + persisted inductor cache. Remaining Phase-2: ratatoskr client cutover.
This commit is contained in:
@@ -0,0 +1,147 @@
|
||||
"""OpenAI-compatible /v1/audio/speech server over dots.tts (rednote-hilab).
|
||||
|
||||
Thin wrapper around DotsTtsRuntime — chosen over SGLang Omni because Omni's edge
|
||||
(continuous batching) is MeanFlow-only and unneeded for a single-consumer surface,
|
||||
while the raw runtime with optimize=True already streams at RTF ~0.22 on our 3090.
|
||||
|
||||
Voice registry: every <name>.wav (+ optional <name>.txt transcript) under
|
||||
DOTS_VOICES_DIR becomes a callable voice. dots.tts REQUIRES an accurate,
|
||||
sentence-bounded transcript to clone cleanly (see the voices/ corpus) — the .txt
|
||||
is that transcript; without it the model leaks reference audio into the output.
|
||||
"""
|
||||
import io
|
||||
import os
|
||||
import glob
|
||||
import struct
|
||||
import threading
|
||||
import wave
|
||||
|
||||
import numpy as np
|
||||
from fastapi import FastAPI, HTTPException
|
||||
from fastapi.responses import Response, StreamingResponse
|
||||
from pydantic import BaseModel
|
||||
|
||||
from dots_tts.runtime import DotsTtsRuntime
|
||||
|
||||
MODEL = os.environ.get("DOTS_MODEL", "dots-studio/dots.tts-soar")
|
||||
VOICES_DIR = os.environ.get("DOTS_VOICES_DIR", "/voices")
|
||||
DEFAULT_VOICE = os.environ.get("DOTS_DEFAULT_VOICE", "donut")
|
||||
NUM_STEPS = int(os.environ.get("DOTS_NUM_STEPS", "10"))
|
||||
GUIDANCE = float(os.environ.get("DOTS_GUIDANCE_SCALE", "1.2"))
|
||||
SAMPLE_RATE = 48000 # dots.tts fixed native output
|
||||
|
||||
app = FastAPI(title="dots.tts")
|
||||
_rt = None
|
||||
_voices: dict = {}
|
||||
# One DotsTtsRuntime, and it is NOT safe to call concurrently (CUDA-graph capture
|
||||
# + shared state). uvicorn runs sync endpoints in a threadpool, so we must
|
||||
# serialize generation ourselves: requests queue and run one at a time. This is
|
||||
# the deliberate trade for the thin-wrapper design — no vLLM-style continuous
|
||||
# batching. If concurrency demand appears, swap the backend to SGLang Omni + the
|
||||
# mf variant behind this same API (see README).
|
||||
_gen_lock = threading.Lock()
|
||||
|
||||
|
||||
def _load_voices() -> dict:
|
||||
reg = {}
|
||||
for wav in sorted(glob.glob(os.path.join(VOICES_DIR, "*.wav"))):
|
||||
name = os.path.splitext(os.path.basename(wav))[0]
|
||||
txt = os.path.splitext(wav)[0] + ".txt"
|
||||
reg[name] = {
|
||||
"wav": wav,
|
||||
"text": open(txt).read().strip() if os.path.exists(txt) else "",
|
||||
}
|
||||
return reg
|
||||
|
||||
|
||||
@app.on_event("startup")
|
||||
def _startup():
|
||||
global _rt, _voices
|
||||
_voices = _load_voices()
|
||||
_rt = DotsTtsRuntime.from_pretrained(MODEL, precision="bfloat16", optimize=True)
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
def health():
|
||||
return {
|
||||
"status": "ok" if _rt is not None else "loading",
|
||||
"model": MODEL,
|
||||
"sample_rate": SAMPLE_RATE,
|
||||
"voices": sorted(_voices),
|
||||
}
|
||||
|
||||
|
||||
@app.get("/v1/voices")
|
||||
def list_voices():
|
||||
return {"voices": sorted(_voices)}
|
||||
|
||||
|
||||
class SpeechRequest(BaseModel):
|
||||
input: str
|
||||
voice: str = DEFAULT_VOICE
|
||||
model: str | None = None # accepted, ignored (single served model)
|
||||
response_format: str = "wav" # wav | pcm
|
||||
stream: bool = False
|
||||
|
||||
|
||||
def _to_pcm16(audio: np.ndarray) -> bytes:
|
||||
return np.round(np.clip(audio, -1.0, 1.0) * 32767.0).astype("<i2").tobytes()
|
||||
|
||||
|
||||
def _wav_bytes(pcm: bytes) -> bytes:
|
||||
buf = io.BytesIO()
|
||||
w = wave.open(buf, "wb")
|
||||
w.setnchannels(1)
|
||||
w.setsampwidth(2)
|
||||
w.setframerate(SAMPLE_RATE)
|
||||
w.writeframes(pcm)
|
||||
w.close()
|
||||
return buf.getvalue()
|
||||
|
||||
|
||||
def _streaming_wav_header() -> bytes:
|
||||
"""WAV header with placeholder (max) sizes — lets a client start playing the
|
||||
stream before the total length is known (the pattern the Zonos/chatterbox
|
||||
consumers already expect)."""
|
||||
return (
|
||||
b"RIFF" + struct.pack("<I", 0xFFFFFFFF) + b"WAVE"
|
||||
+ b"fmt " + struct.pack("<IHHIIHH", 16, 1, 1, SAMPLE_RATE, SAMPLE_RATE * 2, 2, 16)
|
||||
+ b"data" + struct.pack("<I", 0xFFFFFFFF)
|
||||
)
|
||||
|
||||
|
||||
@app.post("/v1/audio/speech")
|
||||
def speech(req: SpeechRequest):
|
||||
rt = _rt
|
||||
if rt is None:
|
||||
raise HTTPException(503, "model still loading")
|
||||
if req.voice not in _voices:
|
||||
raise HTTPException(404, f"unknown voice '{req.voice}'; have {sorted(_voices)}")
|
||||
if not req.input.strip():
|
||||
raise HTTPException(400, "empty input")
|
||||
|
||||
v = _voices[req.voice]
|
||||
kw = dict(
|
||||
prompt_audio_path=v["wav"],
|
||||
prompt_text=v["text"],
|
||||
num_steps=NUM_STEPS,
|
||||
guidance_scale=GUIDANCE,
|
||||
normalize_text=True,
|
||||
)
|
||||
|
||||
if req.stream:
|
||||
def gen():
|
||||
# Hold the lock for the whole stream — a second generation on the
|
||||
# shared runtime mid-stream would corrupt both.
|
||||
with _gen_lock:
|
||||
yield _streaming_wav_header()
|
||||
for chunk in rt.generate_stream(text=req.input, **kw):
|
||||
yield _to_pcm16(chunk.float().cpu().squeeze().numpy())
|
||||
return StreamingResponse(gen(), media_type="audio/wav")
|
||||
|
||||
with _gen_lock:
|
||||
res = rt.generate(text=req.input, **kw)
|
||||
pcm = _to_pcm16(res["audio"].float().cpu().squeeze().numpy())
|
||||
if req.response_format == "pcm":
|
||||
return Response(pcm, media_type="audio/L16;rate=48000")
|
||||
return Response(_wav_bytes(pcm), media_type="audio/wav")
|
||||
Reference in New Issue
Block a user