diff --git a/Makefile b/Makefile index 9d2cb52..5790341 100644 --- a/Makefile +++ b/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\")" diff --git a/docker-compose.yml b/docker-compose.yml index f6609d2..08b07b5 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -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: diff --git a/llama-server.Dockerfile b/llama-server.Dockerfile new file mode 100644 index 0000000..247d974 --- /dev/null +++ b/llama-server.Dockerfile @@ -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"] diff --git a/pyproject.toml b/pyproject.toml index a37c88a..6936422 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", diff --git a/src/tts_mcp/__main__.py b/src/tts_mcp/__main__.py index 2315adb..27ec9e1 100644 --- a/src/tts_mcp/__main__.py +++ b/src/tts_mcp/__main__.py @@ -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__": diff --git a/src/tts_mcp/engines/orpheus.py b/src/tts_mcp/engines/orpheus.py index 462f3f3..69968fa 100644 --- a/src/tts_mcp/engines/orpheus.py +++ b/src/tts_mcp/engines/orpheus.py @@ -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 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 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"") +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"") + 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 strings - token_strings = TOKEN_PATTERN.findall(resp_text) - token_strings = [f"" 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)} diff --git a/src/tts_mcp/server.py b/src/tts_mcp/server.py index ed4f55f..6a31531 100644 --- a/src/tts_mcp/server.py +++ b/src/tts_mcp/server.py @@ -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 etc.)." + "orpheus (LLM via llama-server, supports etc.)." ), lifespan=app_lifespan, ) diff --git a/src/tts_mcp/settings.py b/src/tts_mcp/settings.py index 3933a22..24f116e 100644 --- a/src/tts_mcp/settings.py +++ b/src/tts_mcp/settings.py @@ -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" diff --git a/uv.lock b/uv.lock index 727605f..74f72b5 100644 --- a/uv.lock +++ b/uv.lock @@ -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" },