stacks/index-tts: add streaming WAV endpoint (wrapper 0.2.0)

IndexTTS-2's tts.infer(stream_return=True) is a generator that yields
audio chunks per text segment as they finish, plus inter-segment
silence. Expose this via the existing POST /v1/audio/speech with a new
"stream": true field on the request body.

Wire-up:
  - 44-byte WAV header emitted up front with placeholder data length
    (0xFFFFFFFF) so chunks can be written before total samples are
    known. Players that read until EOF (mpv, ffplay, aplay, sox,
    browsers via <audio>) handle this fine.
  - Each yielded chunk goes through _chunk_to_pcm_bytes(), which
    handles torch tensors / numpy arrays in either int16 or float
    (-1..1) form.
  - 22050 Hz mono int16 — IndexTTS-2's hardcoded output shape.

Time-to-first-audio drops from full-file latency to ~one-segment
latency. Single-sentence inputs barely benefit; long passages /
multi-paragraph reads benefit a lot. Strict metadata parsers may
balk at the placeholder size — request without stream for a
closed-length WAV in that case.

INDEX_TTS_TAG bumped to v2 to force a rebuild.
This commit is contained in:
vh
2026-04-25 14:50:29 -07:00
parent ab696ecbd1
commit 54fef0e9d8
3 changed files with 115 additions and 12 deletions
+81 -9
View File
@@ -30,7 +30,7 @@ import io
import logging
import os
from pathlib import Path
from typing import List, Optional
from typing import Iterable, List, Optional
# IndexTTS pins HF_HUB_CACHE at import time (./checkpoints/hf_cache by
# default, see infer_v2.py:4). Override BEFORE the indextts import or
@@ -40,13 +40,19 @@ os.environ.setdefault(
os.environ.get("INDEX_TTS_HF_CACHE", "/app/checkpoints/hf_cache"),
)
import numpy as np # noqa: E402
import soundfile as sf # noqa: E402
from fastapi import FastAPI, HTTPException # noqa: E402
from fastapi.responses import Response # noqa: E402
from fastapi.responses import Response, StreamingResponse # noqa: E402
from pydantic import BaseModel, Field # noqa: E402
from indextts.infer_v2 import IndexTTS2 # noqa: E402
# IndexTTS-2's hardcoded output rate (infer_v2.py:527). Mono int16.
SR = 22050
CHANNELS = 1
BPS = 16
# ── config from env ──────────────────────────────────────────────────
MODEL_DIR = os.environ.get("INDEX_TTS_MODEL_DIR", "/app/checkpoints")
CFG_PATH = os.environ.get("INDEX_TTS_CFG", f"{MODEL_DIR}/config.yaml")
@@ -73,7 +79,7 @@ tts = IndexTTS2(
)
log.info("IndexTTS2 ready")
app = FastAPI(title="index-tts", version="0.1.0")
app = FastAPI(title="index-tts", version="0.2.0")
class SpeechRequest(BaseModel):
@@ -81,6 +87,17 @@ class SpeechRequest(BaseModel):
input: str = Field(..., description="Text to synthesize")
voice: str = Field(..., description="<name>.wav must exist in voices dir")
response_format: str = Field("wav", description="wav (only)")
stream: bool = Field(
False,
description=(
"If true, stream the WAV as it generates. Each text segment "
"(~120 tokens) yields a chunk as soon as IndexTTS-2 finishes "
"synthesizing it; inter-segment silence is yielded between. "
"Time-to-first-audio drops dramatically for long inputs. The "
"WAV header carries placeholder data size (0xFFFFFFFF) so most "
"players read until EOF."
),
)
# ── emotion (all optional, mutually exclusive) ──
emotion_voice: Optional[str] = Field(
None, description="<name>.wav in emotions dir, used as emotion ref"
@@ -95,6 +112,47 @@ class SpeechRequest(BaseModel):
emotion_alpha: float = Field(1.0, ge=0.0, le=1.0)
def _wav_header(sample_rate: int = SR, channels: int = CHANNELS,
bits_per_sample: int = BPS) -> bytes:
"""44-byte RIFF/WAVE/PCM header with placeholder data length so the
payload can be streamed without knowing total samples up front.
Players that read until EOF (mpv, ffplay, aplay, sox, browsers) handle
this fine. Strict parsers (some metadata extractors) may complain."""
byte_rate = sample_rate * channels * bits_per_sample // 8
block_align = channels * bits_per_sample // 8
placeholder = 0xFFFFFFFF
return (
b"RIFF"
+ placeholder.to_bytes(4, "little")
+ b"WAVE"
+ b"fmt "
+ (16).to_bytes(4, "little")
+ (1).to_bytes(2, "little") # PCM
+ channels.to_bytes(2, "little")
+ sample_rate.to_bytes(4, "little")
+ byte_rate.to_bytes(4, "little")
+ block_align.to_bytes(2, "little")
+ bits_per_sample.to_bytes(2, "little")
+ b"data"
+ placeholder.to_bytes(4, "little")
)
def _chunk_to_pcm_bytes(chunk) -> bytes:
"""Normalize whatever IndexTTS-2 yields (torch tensor, numpy array,
int16 or float) into raw little-endian int16 PCM bytes."""
if hasattr(chunk, "cpu"): # torch.Tensor
chunk = chunk.cpu().numpy()
arr = np.asarray(chunk).reshape(-1) # flatten to mono samples
if arr.dtype != np.int16:
# If the model yields float (-1..1), scale into int16 range.
if np.issubdtype(arr.dtype, np.floating):
arr = np.clip(arr * 32767.0, -32768, 32767).astype(np.int16)
else:
arr = arr.astype(np.int16)
return arr.tobytes()
def _resolve(name: str, root: Path) -> Path:
p = root / f"{name}.wav"
if not p.is_file():
@@ -116,7 +174,7 @@ def list_voices() -> dict:
@app.post("/v1/audio/speech")
def synthesize(req: SpeechRequest) -> Response:
def synthesize(req: SpeechRequest):
if req.response_format != "wav":
raise HTTPException(status_code=400, detail="only response_format=wav is supported")
@@ -137,10 +195,10 @@ def synthesize(req: SpeechRequest) -> Response:
use_emo_text = True
emo_text = req.emotion_text
sr, audio = tts.infer(
common_kwargs = dict(
spk_audio_prompt=spk,
text=req.input,
output_path=None, # in-memory return: (sr, np_int16)
output_path=None,
emo_audio_prompt=emo_path,
emo_alpha=req.emotion_alpha,
emo_vector=emo_vector,
@@ -149,6 +207,20 @@ def synthesize(req: SpeechRequest) -> Response:
verbose=False,
)
buf = io.BytesIO()
sf.write(buf, audio, sr, format="WAV", subtype="PCM_16")
return Response(content=buf.getvalue(), media_type="audio/wav")
if not req.stream:
sr, audio = tts.infer(**common_kwargs)
buf = io.BytesIO()
sf.write(buf, audio, sr, format="WAV", subtype="PCM_16")
return Response(content=buf.getvalue(), media_type="audio/wav")
# Streaming path. tts.infer with stream_return=True is a generator
# yielding torch audio tensors per text segment plus inter-segment
# silence. We write a streaming-friendly WAV header up front, then
# int16 PCM bytes per chunk. Sample rate / channels / bps are fixed
# by IndexTTS-2 (22050 Hz mono 16-bit).
def iter_wav() -> Iterable[bytes]:
yield _wav_header()
for chunk in tts.infer(stream_return=True, **common_kwargs):
yield _chunk_to_pcm_bytes(chunk)
return StreamingResponse(iter_wav(), media_type="audio/wav")