kokoro-tts stream added

This commit is contained in:
2026-10-03 15:20:33 +03:00
parent 509fd3dde0
commit 1c77ae658e
4 changed files with 213 additions and 46 deletions
+171 -16
View File
@@ -10,10 +10,11 @@ non-native (measured 27% vs 22% round-trip WER, per the model card).
So: text -> RuG2P.phonemize -> KModel(ipa, voicepack[len(ipa) - 1]) -> waveform.
Endpoints
POST /v1/audio/speech OpenAI text-to-speech
GET /v1/models OpenAI model list
GET /v1/voices voice inventory (extension, not part of OpenAI)
GET /healthz readiness, 503 until the model is loaded
POST /v1/audio/speech OpenAI text-to-speech
POST /v1/audio/speech/stream same, but mp3/opus emitted while synthesising
GET /v1/models OpenAI model list
GET /v1/voices voice inventory (extension, not part of OpenAI)
GET /healthz readiness, 503 until the model is loaded
"""
from __future__ import annotations
@@ -21,6 +22,7 @@ from __future__ import annotations
import io
import logging
import os
import queue
import re
import subprocess
import sys
@@ -28,11 +30,11 @@ import threading
import wave
from contextlib import asynccontextmanager
from pathlib import Path
from typing import TYPE_CHECKING, Literal
from typing import TYPE_CHECKING, Iterator, Literal
import numpy as np
from fastapi import FastAPI
from fastapi.responses import JSONResponse, Response
from fastapi.responses import JSONResponse, Response, StreamingResponse
from pydantic import BaseModel, ConfigDict, Field
if TYPE_CHECKING: # torch is imported lazily so /healthz answers during boot
@@ -42,7 +44,12 @@ MODEL_ID = "kokoro-ru"
SAMPLE_RATE = 24000
MODEL_DIR = Path(os.environ.get("KOKORO_MODEL_DIR", "/app/kokoro-ru"))
DEFAULT_VOICE = os.environ.get("KOKORO_DEFAULT_VOICE", "sveta")
THREADS = int(os.environ.get("KOKORO_THREADS", os.cpu_count() or 4))
# Measured on the host this was tuned for (Ryzen AI 9 HX 370, 24 logical cores):
# median end-to-end latency for a 5.6 s utterance was 1.203 s @ 4 threads,
# 1.066 s @ 8, 0.979 s @ 12, 0.980 s @ 16, then 1.87 s @ 24. The gain stops at
# the physical core count and SMT oversubscription costs ~2x, so cap instead of
# trusting os.cpu_count(), which reports logical CPUs. Override on other hosts.
THREADS = int(os.environ.get("KOKORO_THREADS", min(12, os.cpu_count() or 4)))
# 2026-07-29, when the kokoro-ru revision we pin was published. Clients that
# cache on this treat any change as a new model, so it must stay stable.
MODEL_CREATED = 1785353253
@@ -223,30 +230,39 @@ class KokoroRu:
if ps:
yield from split_phonemes(ps)
def synthesize(self, text: str, voice: str, speed: float) -> np.ndarray:
def iter_audio_chunks(self, text: str, voice: str, speed: float):
"""Yields float32 audio per phoneme chunk, silence gaps interleaved.
The engine lock is held for the whole iteration, so a caller that stops
consuming early releases synthesis for everyone else.
"""
torch = self._torch
assert torch is not None, "synthesize() before load()"
stem, _gender = VOICE_SPECS[voice]
model = self._models[stem]
pack = self._packs[voice]
gap = torch.zeros(int(CHUNK_GAP_S * SAMPLE_RATE), dtype=torch.float32)
pieces: list[torch.Tensor] = []
gap = np.zeros(int(CHUNK_GAP_S * SAMPLE_RATE), dtype=np.float32)
with self._lock:
for ps in self.phonemes(text):
for index, ps in enumerate(self.phonemes(text)):
# The style vector is picked by phoneme-string length, which is
# why the model sounds deterministic for identical text.
style = pack[len(ps) - 1]
# The packs ship as [510, 256]; KModel wants a batch of one.
if style.dim() == 1:
style = style.unsqueeze(0)
if pieces:
pieces.append(gap)
pieces.append(model(ps, style, speed, return_output=True).audio)
if index:
yield gap
yield np.asarray(
model(ps, style, speed, return_output=True).audio,
dtype=np.float32,
).reshape(-1)
if not pieces:
def synthesize(self, text: str, voice: str, speed: float) -> np.ndarray:
chunks = list(self.iter_audio_chunks(text, voice, speed))
if not chunks:
return np.zeros(0, dtype=np.float32)
return torch.cat(pieces).numpy().astype(np.float32, copy=False)
return np.concatenate(chunks)
def encode(audio: np.ndarray, fmt: str) -> bytes:
@@ -290,6 +306,86 @@ def encode(audio: np.ndarray, fmt: str) -> bytes:
return done.stdout
class StreamEncoder:
"""One long-lived ffmpeg per request: raw PCM in, encoded bytes out.
A single process is what keeps the container valid. Handing it the audio in
pieces as they are synthesised avoids any byte-level concatenation, whereas
encoding the pieces separately and joining the results would emit chained
Ogg for opus, which plenty of players reject.
"""
def __init__(self, fmt: str) -> None:
import imageio_ffmpeg
self._proc = subprocess.Popen(
[
imageio_ffmpeg.get_ffmpeg_exe(),
"-hide_banner",
"-loglevel",
"error",
"-f",
"s16le",
"-ar",
str(SAMPLE_RATE),
"-ac",
"1",
"-i",
"pipe:0",
*FFMPEG_ARGS[fmt],
"-f",
FFMPEG_CONTAINERS[fmt],
"pipe:1",
],
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
self._blocks: queue.Queue[bytes | None] = queue.Queue()
self._reader = threading.Thread(target=self._pump, daemon=True)
self._reader.start()
def _pump(self) -> None:
assert self._proc.stdout is not None
while True:
block = self._proc.stdout.read(8192)
if not block:
break
self._blocks.put(block)
self._blocks.put(None)
def push(self, audio: np.ndarray) -> None:
assert self._proc.stdin is not None
clipped = np.clip(audio, -1.0, 1.0)
self._proc.stdin.write((clipped * 32767.0).astype("<i2").tobytes())
self._proc.stdin.flush()
def drain(self) -> Iterator[bytes]:
"""Yields whatever ffmpeg has already emitted, without waiting for more."""
while True:
try:
block = self._blocks.get_nowait()
except queue.Empty:
return
if block is None:
return
yield block
def finish(self) -> Iterator[bytes]:
assert self._proc.stdin is not None
self._proc.stdin.close()
self._reader.join(timeout=120)
code = self._proc.wait(timeout=30)
error = self._proc.stderr.read().decode("utf-8", "replace").strip()[-400:]
if code != 0:
raise RuntimeError(error or f"ffmpeg exited with {code}")
yield from self.drain()
def abort(self) -> None:
if self._proc.poll() is None:
self._proc.kill()
engine = KokoroRu()
state: dict[str, str | None] = {"status": "loading", "error": None}
@@ -331,6 +427,13 @@ class SpeechRequest(BaseModel):
speed: float | None = Field(default=None, ge=0.25, le=4.0)
class StreamSpeechRequest(SpeechRequest):
# Streaming needs a container that tolerates unknown length up front, so wav
# (whose header declares the final sizes) and the raw formats are out. mp3
# and opus emit bytes as they go, which is the whole point of the endpoint.
response_format: Literal["mp3", "opus"] = "mp3"
def fail(status: int, message: str, param: str | None = None, code: str | None = None) -> JSONResponse:
return JSONResponse(
status_code=status,
@@ -395,6 +498,58 @@ def create_speech(request: SpeechRequest) -> Response | JSONResponse:
)
# response_model=None for the same reason as create_speech above.
@app.post("/v1/audio/speech/stream", response_model=None)
def stream_speech(request: StreamSpeechRequest) -> Response | JSONResponse:
if state["status"] != "ready":
return fail(503, f"model is not ready: {state['status']}", code="model_not_ready")
voice = resolve_voice(request.voice)
if voice is None:
available = ", ".join(engine.available_voices())
return fail(
400,
f"unknown voice {request.voice!r}; available: {available}",
param="voice",
code="unknown_voice",
)
chunks = engine.iter_audio_chunks(request.input, voice, request.speed or 1.0)
try:
# Pulled before responding: once the status line is sent it cannot become
# a 400, and input with no speakable text has to keep failing that way.
first = next(chunks)
except StopIteration:
return fail(
400,
"input contains no speakable text for the Russian G2P",
param="input",
code="no_phonemes",
)
def body() -> Iterator[bytes]:
encoder = StreamEncoder(request.response_format)
try:
encoder.push(first)
yield from encoder.drain()
for chunk in chunks:
encoder.push(chunk)
yield from encoder.drain()
yield from encoder.finish()
except Exception:
log.exception("streaming synthesis failed")
raise
finally:
chunks.close()
encoder.abort()
return StreamingResponse(
body(),
media_type=CONTENT_TYPES[request.response_format],
headers={"model-id": MODEL_ID, "voice-id": voice},
)
@app.get("/v1/models")
def list_models() -> dict:
return {