Replace Ollama with llama-server for 15x Orpheus throughput
Build llama.cpp from source with SM 120 CUDA kernels and FORCE_CUBLAS for RTX 5070 Blackwell. Rewrite OrpheusEngine to stream tokens via SSE and decode SNAC in overlapping 28-token batches (4 frames), replacing the blocking requests+stream:false approach. Performance: 13.5 → 170-213 tok/s. 100s audio generates in ~48s (2x faster than realtime). Replaces requests with httpx async client. Also switch MCP transport to stateless_http mode so container restarts don't invalidate client sessions.
This commit is contained in:
parent
7f9557d93b
commit
538b8a513e
9
Makefile
9
Makefile
@ -1,7 +1,7 @@
|
||||
include .env
|
||||
export
|
||||
|
||||
.PHONY: build up down logs restart status
|
||||
.PHONY: build up down logs restart status bench
|
||||
|
||||
build:
|
||||
docker compose build
|
||||
@ -26,3 +26,10 @@ status:
|
||||
@docker compose ps
|
||||
@echo "---"
|
||||
@curl -s http://localhost:8371/mcp 2>/dev/null | head -5 || echo "Server not responding"
|
||||
|
||||
bench:
|
||||
@echo "Benchmarking llama-server throughput..."
|
||||
@docker exec orpheus-llama-server curl -s "http://127.0.0.1:8081/v1/completions" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"prompt":"<|audio|>tara: Hello, how are you doing today?<|eot_id|>","max_tokens":500,"stream":false}' | \
|
||||
python3 -c "import sys,json; d=json.load(sys.stdin); u=d['usage']; t=d.get('timings',{}); print(f\"{u['completion_tokens']} tokens, {t.get('predicted_per_second',0):.1f} tok/s\")"
|
||||
|
||||
@ -7,11 +7,9 @@ services:
|
||||
environment:
|
||||
# Override for Docker networking (container DNS instead of IPs)
|
||||
TTS_PIPER_HOST: piper-tts
|
||||
TTS_OLLAMA_URL: http://host.docker.internal:11434
|
||||
TTS_ORPHEUS_URL: http://llama-server:8081
|
||||
# PipeWire client config
|
||||
XDG_RUNTIME_DIR: /run/user/1000
|
||||
extra_hosts:
|
||||
- "host.docker.internal:host-gateway"
|
||||
volumes:
|
||||
# Kokoro ONNX models (read-only)
|
||||
- ./models:/app/models:ro
|
||||
@ -19,6 +17,9 @@ services:
|
||||
- hf-cache:/home/tts/.cache/huggingface
|
||||
# PipeWire socket for audio playback through host speakers
|
||||
- /run/user/1000/pipewire-0:/run/user/1000/pipewire-0
|
||||
depends_on:
|
||||
llama-server:
|
||||
condition: service_healthy
|
||||
networks:
|
||||
- caddy
|
||||
- dootie-internal
|
||||
@ -26,6 +27,36 @@ services:
|
||||
caddy: voice.l.supported.systems
|
||||
caddy.reverse_proxy: "{{upstreams 8371}}"
|
||||
|
||||
llama-server:
|
||||
build:
|
||||
context: .
|
||||
dockerfile: llama-server.Dockerfile
|
||||
container_name: orpheus-llama-server
|
||||
restart: unless-stopped
|
||||
deploy:
|
||||
resources:
|
||||
reservations:
|
||||
devices:
|
||||
- driver: nvidia
|
||||
count: 1
|
||||
capabilities: [gpu]
|
||||
volumes:
|
||||
# GGUF model file (set ORPHEUS_GGUF_PATH in .env, e.g. from Ollama blob storage)
|
||||
- ${ORPHEUS_GGUF_PATH}:/models/orpheus.gguf:ro
|
||||
command: >-
|
||||
--host 0.0.0.0 --port 8081
|
||||
--model /models/orpheus.gguf
|
||||
--n-gpu-layers 999 --ctx-size 4096
|
||||
--flash-attn --cont-batching
|
||||
networks:
|
||||
- dootie-internal
|
||||
healthcheck:
|
||||
test: ["CMD", "curl", "-sf", "http://127.0.0.1:8081/health"]
|
||||
interval: 15s
|
||||
timeout: 5s
|
||||
start_period: 120s
|
||||
retries: 5
|
||||
|
||||
volumes:
|
||||
hf-cache:
|
||||
|
||||
|
||||
48
llama-server.Dockerfile
Normal file
48
llama-server.Dockerfile
Normal file
@ -0,0 +1,48 @@
|
||||
# llama-server built from source with SM 120 (Blackwell / RTX 5070) CUDA kernels.
|
||||
# Multi-stage: ~8GB devel toolkit stays in builder, runtime image is ~2GB.
|
||||
|
||||
# ── Builder ──────────────────────────────────────────────────────────────
|
||||
FROM nvidia/cuda:12.8.1-devel-ubuntu24.04 AS builder
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
cmake git build-essential curl ca-certificates \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
WORKDIR /build
|
||||
|
||||
# Pin to a release tag for reproducibility
|
||||
ARG LLAMA_CPP_VERSION=b5460
|
||||
RUN git clone --depth 1 --branch ${LLAMA_CPP_VERSION} \
|
||||
https://github.com/ggerganov/llama.cpp.git
|
||||
|
||||
WORKDIR /build/llama.cpp
|
||||
|
||||
# CUDA driver symbols (cuMemCreate, etc.) are resolved at runtime by nvidia-container-runtime.
|
||||
# --allow-shlib-undefined lets the linker accept unresolved refs in libggml-cuda.so.
|
||||
RUN cmake -B build \
|
||||
-DCMAKE_BUILD_TYPE=Release \
|
||||
-DCMAKE_CUDA_ARCHITECTURES=120 \
|
||||
-DGGML_CUDA=ON \
|
||||
-DGGML_CUDA_FORCE_CUBLAS=ON \
|
||||
-DLLAMA_BUILD_SERVER=ON \
|
||||
-DLLAMA_CURL=OFF \
|
||||
-DCMAKE_EXE_LINKER_FLAGS="-Wl,--allow-shlib-undefined" \
|
||||
&& cmake --build build --target llama-server -j$(nproc)
|
||||
|
||||
# ── Runtime ──────────────────────────────────────────────────────────────
|
||||
FROM nvidia/cuda:12.8.1-runtime-ubuntu24.04
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
curl ca-certificates libgomp1 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Copy server binary and its shared libraries (libggml-*.so)
|
||||
COPY --from=builder /build/llama.cpp/build/bin/llama-server /usr/local/bin/llama-server
|
||||
COPY --from=builder /build/llama.cpp/build/bin/lib*.so /usr/local/lib/
|
||||
|
||||
RUN ldconfig
|
||||
|
||||
RUN useradd -u 1000 -m llama 2>/dev/null || true
|
||||
USER 1000
|
||||
|
||||
ENTRYPOINT ["llama-server"]
|
||||
@ -11,11 +11,11 @@ license = "MIT"
|
||||
authors = [{name = "Ryan Malloy", email = "ryan@supported.systems"}]
|
||||
dependencies = [
|
||||
"fastmcp>=3.0.0",
|
||||
"httpx",
|
||||
"kokoro-onnx>=0.5.0",
|
||||
"numpy",
|
||||
"onnxruntime",
|
||||
"pydantic-settings",
|
||||
"requests",
|
||||
"snac>=1.2.1",
|
||||
"soundfile",
|
||||
"torch",
|
||||
|
||||
@ -5,7 +5,12 @@ from .settings import settings
|
||||
|
||||
|
||||
def main():
|
||||
mcp.run(transport="streamable-http", host=settings.host, port=settings.port)
|
||||
mcp.run(
|
||||
transport="streamable-http",
|
||||
host=settings.host,
|
||||
port=settings.port,
|
||||
stateless_http=True,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@ -1,18 +1,20 @@
|
||||
"""Orpheus TTS via Ollama + SNAC decoder.
|
||||
"""Orpheus TTS via llama-server + streaming SNAC decoder.
|
||||
|
||||
Sends text to Ollama's Orpheus model, parses <custom_token_N> responses,
|
||||
and decodes through SNAC to 24kHz WAV. SNAC is lazy-loaded on first use
|
||||
and runs on CPU (RTX 5070 SM 120 not yet supported by PyTorch 2.6).
|
||||
Sends text to a llama-server completions endpoint with streaming enabled,
|
||||
parses <custom_token_N> responses as they arrive via SSE, and decodes
|
||||
through SNAC in batches of 28 tokens (4 frames) for overlapped inference.
|
||||
SNAC runs on CPU; the LLM runs on GPU via llama-server.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
import time
|
||||
|
||||
import httpx
|
||||
import numpy as np
|
||||
import requests
|
||||
|
||||
from ..audio import wav_duration, write_wav
|
||||
from ..settings import settings
|
||||
@ -21,6 +23,9 @@ from .base import TTSEngine, TTSResult
|
||||
SAMPLE_RATE = 24000
|
||||
SNAC_CODEBOOK_SIZE = 4096
|
||||
TOKEN_PATTERN = re.compile(r"<custom_token_(\d+)>")
|
||||
TOKENS_PER_FRAME = 7
|
||||
FRAMES_PER_BATCH = 4
|
||||
BATCH_SIZE = TOKENS_PER_FRAME * FRAMES_PER_BATCH # 28
|
||||
|
||||
ALL_VOICES = ["tara", "leah", "jess", "leo", "dan", "mia", "zac", "zoe"]
|
||||
|
||||
@ -87,7 +92,12 @@ def _tokens_to_audio(token_strings: list[str], snac_model) -> np.ndarray | None:
|
||||
|
||||
|
||||
def _load_snac():
|
||||
"""Load SNAC model on CPU. Called once, lazily."""
|
||||
"""Load SNAC model on CPU. Called once, lazily.
|
||||
|
||||
Sets CUDA_VISIBLE_DEVICES="" for the entire process to prevent PyTorch
|
||||
from allocating GPU memory for SNAC. Safe because all GPU inference is
|
||||
handled by llama-server in a separate container.
|
||||
"""
|
||||
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
||||
from snac import SNAC
|
||||
|
||||
@ -95,20 +105,28 @@ def _load_snac():
|
||||
|
||||
|
||||
class OrpheusEngine(TTSEngine):
|
||||
"""Orpheus TTS via Ollama's completions API + SNAC audio decoding.
|
||||
"""Orpheus TTS via llama-server completions API + streaming SNAC decode.
|
||||
|
||||
SNAC is lazy-loaded on first synthesize() call to avoid holding
|
||||
~200MB of RAM when the engine isn't being used.
|
||||
Streams tokens from llama-server via SSE, decodes in batches of 28
|
||||
tokens (4 SNAC frames) to overlap GPU inference with CPU audio decode.
|
||||
SNAC is lazy-loaded on first synthesize() call.
|
||||
"""
|
||||
|
||||
name = "orpheus"
|
||||
default_voice = "tara"
|
||||
|
||||
def __init__(self, ollama_url: str, model_name: str) -> None:
|
||||
def __init__(self, orpheus_url: str) -> None:
|
||||
self._snac = None
|
||||
self._snac_lock = asyncio.Lock()
|
||||
self._ollama_url = ollama_url
|
||||
self._model = model_name
|
||||
self._url = orpheus_url
|
||||
self._client = httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(connect=10.0, read=600.0, write=30.0, pool=30.0),
|
||||
limits=httpx.Limits(max_connections=4, max_keepalive_connections=2),
|
||||
)
|
||||
|
||||
async def close(self) -> None:
|
||||
"""Clean up httpx client."""
|
||||
await self._client.aclose()
|
||||
|
||||
async def _get_snac(self):
|
||||
"""Lazy-load SNAC on first use."""
|
||||
@ -129,66 +147,119 @@ class OrpheusEngine(TTSEngine):
|
||||
prompt = f"<|audio|>{voice}: {text}<|eot_id|>"
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
snac = await self._get_snac()
|
||||
|
||||
# Ollama API call (blocking HTTP)
|
||||
# Stream tokens from llama-server via SSE
|
||||
t0 = time.time()
|
||||
token_strings: list[str] = []
|
||||
audio_chunks: list[np.ndarray] = []
|
||||
total_tokens = 0
|
||||
dropped_lines = 0
|
||||
stream_interrupted = False
|
||||
|
||||
def _call_ollama():
|
||||
resp = requests.post(
|
||||
f"{self._ollama_url}/v1/completions",
|
||||
try:
|
||||
async with self._client.stream(
|
||||
"POST",
|
||||
f"{self._url}/v1/completions",
|
||||
json={
|
||||
"model": self._model,
|
||||
"prompt": prompt,
|
||||
"max_tokens": 8192,
|
||||
"temperature": 0.6,
|
||||
"top_p": 0.9,
|
||||
"stream": False,
|
||||
"stream": True,
|
||||
},
|
||||
timeout=600,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
) as response:
|
||||
try:
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as e:
|
||||
body = await response.aread()
|
||||
raise RuntimeError(
|
||||
f"llama-server returned {e.response.status_code}: "
|
||||
f"{body[:500].decode(errors='replace')}"
|
||||
) from e
|
||||
|
||||
async for line in response.aiter_lines():
|
||||
# SSE format: "data: {...}" or "data: [DONE]"
|
||||
if not line.startswith("data: "):
|
||||
continue
|
||||
payload = line[6:]
|
||||
if payload == "[DONE]":
|
||||
break
|
||||
|
||||
try:
|
||||
chunk = json.loads(payload)
|
||||
except json.JSONDecodeError:
|
||||
dropped_lines += 1
|
||||
continue
|
||||
|
||||
chunk_text = chunk.get("choices", [{}])[0].get("text", "")
|
||||
matches = TOKEN_PATTERN.findall(chunk_text)
|
||||
for m in matches:
|
||||
token_val = int(m)
|
||||
# Skip leading special tokens (value < 10)
|
||||
if total_tokens == 0 and token_val < 10:
|
||||
continue
|
||||
token_strings.append(f"<custom_token_{m}>")
|
||||
total_tokens += 1
|
||||
|
||||
# Batch decode every 28 tokens (4 SNAC frames)
|
||||
while len(token_strings) >= BATCH_SIZE:
|
||||
batch = token_strings[:BATCH_SIZE]
|
||||
token_strings = token_strings[BATCH_SIZE:]
|
||||
chunk_audio = await loop.run_in_executor(
|
||||
None, _tokens_to_audio, batch, snac
|
||||
)
|
||||
if chunk_audio is not None:
|
||||
audio_chunks.append(chunk_audio)
|
||||
|
||||
except httpx.ConnectError as e:
|
||||
raise RuntimeError(f"Cannot reach llama-server at {self._url}: {e}") from e
|
||||
except (httpx.ReadError, httpx.RemoteProtocolError) as e:
|
||||
if not audio_chunks:
|
||||
raise RuntimeError(
|
||||
f"llama-server connection lost with no audio decoded: {e}"
|
||||
) from e
|
||||
stream_interrupted = True
|
||||
print(
|
||||
f" Warning: llama-server connection lost after {total_tokens} tokens. "
|
||||
f"Using {len(audio_chunks)} partial chunks.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
result = await loop.run_in_executor(None, _call_ollama)
|
||||
gen_time = time.time() - t0
|
||||
|
||||
resp_text = result.get("choices", [{}])[0].get("text", "")
|
||||
if dropped_lines > 0:
|
||||
print(
|
||||
f" Warning: {dropped_lines} SSE lines had malformed JSON",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
# Extract <custom_token_N> strings
|
||||
token_strings = TOKEN_PATTERN.findall(resp_text)
|
||||
token_strings = [f"<custom_token_{t}>" for t in token_strings]
|
||||
|
||||
# Skip leading special tokens (values < 10)
|
||||
skip = 0
|
||||
for ts in token_strings:
|
||||
m = TOKEN_PATTERN.search(ts)
|
||||
if m and int(m.group(1)) < 10:
|
||||
skip += 1
|
||||
else:
|
||||
break
|
||||
if skip > 0:
|
||||
token_strings = token_strings[skip:]
|
||||
# Flush remaining tokens (drop partial frame — at most 0.29ms lost)
|
||||
if token_strings:
|
||||
usable = len(token_strings) - (len(token_strings) % TOKENS_PER_FRAME)
|
||||
if usable > 0:
|
||||
chunk_audio = await loop.run_in_executor(
|
||||
None, _tokens_to_audio, token_strings[:usable], snac
|
||||
)
|
||||
if chunk_audio is not None:
|
||||
audio_chunks.append(chunk_audio)
|
||||
|
||||
num_frames = total_tokens // TOKENS_PER_FRAME
|
||||
tok_per_sec = total_tokens / gen_time if gen_time > 0 else 0
|
||||
status = " (TRUNCATED)" if stream_interrupted else ""
|
||||
print(
|
||||
f" Orpheus: {len(token_strings)} tokens "
|
||||
f"({len(token_strings) // 7} frames) in {gen_time:.1f}s",
|
||||
f" Orpheus: {total_tokens} tokens ({num_frames} frames) "
|
||||
f"in {gen_time:.1f}s ({tok_per_sec:.1f} tok/s){status}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
|
||||
if len(token_strings) < 7:
|
||||
if not audio_chunks:
|
||||
raise RuntimeError(
|
||||
f"Orpheus returned insufficient tokens ({len(token_strings)}). "
|
||||
f"Response preview: {resp_text[:200]}"
|
||||
f"Orpheus returned insufficient tokens ({total_tokens}). "
|
||||
"Check llama-server logs."
|
||||
)
|
||||
|
||||
# Lazy-load SNAC, then decode (CPU-bound)
|
||||
snac = await self._get_snac()
|
||||
audio = await loop.run_in_executor(
|
||||
None, _tokens_to_audio, token_strings, snac
|
||||
)
|
||||
if audio is None:
|
||||
raise RuntimeError("SNAC decoding produced no audio")
|
||||
|
||||
audio = np.concatenate(audio_chunks)
|
||||
path = write_wav(audio, SAMPLE_RATE, prefix="orpheus-")
|
||||
|
||||
return TTSResult(
|
||||
@ -205,15 +276,15 @@ class OrpheusEngine(TTSEngine):
|
||||
|
||||
async def check_health(self) -> dict:
|
||||
try:
|
||||
resp = requests.get(f"{self._ollama_url}/api/tags", timeout=5)
|
||||
resp = await self._client.get(f"{self._url}/health")
|
||||
resp.raise_for_status()
|
||||
models = [m["name"] for m in resp.json().get("models", [])]
|
||||
has_orpheus = any("orpheus" in m.lower() for m in models)
|
||||
data = resp.json()
|
||||
status = data.get("status", "unknown")
|
||||
return {
|
||||
"status": "healthy" if has_orpheus else "degraded",
|
||||
"status": "healthy" if status == "ok" else "degraded",
|
||||
"engine": self.name,
|
||||
"model_loaded": has_orpheus,
|
||||
"ollama_models": len(models),
|
||||
"backend": "llama-server",
|
||||
"server_status": status,
|
||||
}
|
||||
except Exception as e:
|
||||
return {"status": "unhealthy", "engine": self.name, "error": str(e)}
|
||||
|
||||
@ -44,7 +44,7 @@ async def app_lifespan(server: FastMCP):
|
||||
engines: dict[str, TTSEngine] = {
|
||||
"piper": PiperEngine(settings.piper_host, settings.piper_port),
|
||||
"kokoro": KokoroEngine(kokoro_model),
|
||||
"orpheus": OrpheusEngine(settings.ollama_url, settings.orpheus_model),
|
||||
"orpheus": OrpheusEngine(settings.orpheus_url),
|
||||
}
|
||||
|
||||
# Health check all engines at startup
|
||||
@ -66,6 +66,9 @@ async def app_lifespan(server: FastMCP):
|
||||
finally:
|
||||
print("TTS MCP server shutting down", file=sys.stderr)
|
||||
await queue.stop()
|
||||
for eng in engines.values():
|
||||
if hasattr(eng, "close"):
|
||||
await eng.close()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@ -79,7 +82,7 @@ mcp = FastMCP(
|
||||
"through the host speakers (queued so agents don't talk over each other). "
|
||||
"Use 'generate_audio' to synthesize without playing. "
|
||||
"Engines: kokoro (fast ONNX, ~50 voices), piper (Wyoming/Docker), "
|
||||
"orpheus (Ollama LLM, supports <laugh> etc.)."
|
||||
"orpheus (LLM via llama-server, supports <laugh> etc.)."
|
||||
),
|
||||
lifespan=app_lifespan,
|
||||
)
|
||||
|
||||
@ -20,9 +20,8 @@ class Settings(BaseSettings):
|
||||
kokoro_model: Path = Path("models/kokoro/kokoro-v1.0.onnx")
|
||||
kokoro_voices: Path = Path("models/kokoro/voices-v1.0.bin")
|
||||
|
||||
# Orpheus (Ollama + SNAC)
|
||||
ollama_url: str = "http://127.0.0.1:11434"
|
||||
orpheus_model: str = "legraphista/Orpheus:3b-ft-q4_k_m"
|
||||
# Orpheus (llama-server + SNAC)
|
||||
orpheus_url: str = "http://127.0.0.1:8081"
|
||||
|
||||
# Voice filtering
|
||||
voice_blacklist: str = "amy,jess,zoe,adam"
|
||||
|
||||
4
uv.lock
generated
4
uv.lock
generated
@ -1977,11 +1977,11 @@ version = "2026.2.20"
|
||||
source = { editable = "." }
|
||||
dependencies = [
|
||||
{ name = "fastmcp" },
|
||||
{ name = "httpx" },
|
||||
{ name = "kokoro-onnx" },
|
||||
{ name = "numpy" },
|
||||
{ name = "onnxruntime" },
|
||||
{ name = "pydantic-settings" },
|
||||
{ name = "requests" },
|
||||
{ name = "snac" },
|
||||
{ name = "soundfile" },
|
||||
{ name = "torch" },
|
||||
@ -1991,11 +1991,11 @@ dependencies = [
|
||||
[package.metadata]
|
||||
requires-dist = [
|
||||
{ name = "fastmcp", specifier = ">=3.0.0" },
|
||||
{ name = "httpx" },
|
||||
{ name = "kokoro-onnx", specifier = ">=0.5.0" },
|
||||
{ name = "numpy" },
|
||||
{ name = "onnxruntime" },
|
||||
{ name = "pydantic-settings" },
|
||||
{ name = "requests" },
|
||||
{ name = "snac", specifier = ">=1.2.1" },
|
||||
{ name = "soundfile" },
|
||||
{ name = "torch" },
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user