feat(morpheus): permanent mOrpheus TTS stack (vLLM bf16 + SNAC/FastAPI wrapper) on irv-ml1

Two-container stack serving MrDragonFox/mOrpheus (uncensored Orpheus TTS, Llama-3.2-3B
-> SNAC 24kHz). vllm-morpheus (GPU/3090) emits Orpheus audio tokens; morpheus-tts (CPU)
SNAC-decodes them to WAV and exposes POST /tts (baddy voice + zero-shot cloning). Deployed
+ tested end-to-end (28/28 valid frames, valid WAV, reachable over WG).

Hard-won config, all encoded in compose/README:
- bf16 REQUIRED: --quantization fp8 destroys audio-token generation (0 valid SNAC frames
  even at greedy). Footprint ~7.9GB.
- Image PINNED to v0.23.0: 'latest' ships Blackwell oink/aiter kernels that crash on Ampere
  import.
- 3090 (not the comfy-contended A6000); --enforce-eager to fit the shared card.
- RTF ~1.0 end-to-end (gen ~98 tok/s / RTF 0.84 + CPU decode + HTTP).

INTERNAL RESEARCH ONLY (CC-BY-NC-4.0); do not expose externally.
This commit is contained in:
vh
2026-07-09 00:49:25 -07:00
parent 99a4a1721f
commit 01eedd8d27
6 changed files with 315 additions and 0 deletions
+14
View File
@@ -0,0 +1,14 @@
# mOrpheus TTS wrapper — CPU-only (SNAC decode + FastAPI /tts). Calls the vLLM engine.
FROM python:3.11-slim
RUN apt-get update && apt-get install -y --no-install-recommends libsndfile1 && rm -rf /var/lib/apt/lists/*
WORKDIR /app
COPY requirements.txt .
# CPU torch (SNAC decode is small; keeps this container off the GPU / out of contention)
RUN pip install --no-cache-dir torch --index-url https://download.pytorch.org/whl/cpu \
&& pip install --no-cache-dir -r requirements.txt
COPY app.py .
EXPOSE 8000
CMD ["uvicorn", "app:app", "--host", "0.0.0.0", "--port", "8000"]
+125
View File
@@ -0,0 +1,125 @@
#!/usr/bin/env python3
"""mOrpheus TTS wrapper — turns the vLLM engine's Orpheus audio tokens into 24kHz WAV.
Architecture: this CPU service builds the Orpheus prompt (as raw token ids), calls the
vLLM engine (which serves the mOrpheus LLM), recovers the generated audio-token ids by
re-tokenizing the returned text (vLLM emits them as `<custom_token_N>` strings when
skip_special_tokens=false), then SNAC-decodes them to audio.
Token scheme (verified): audio-base 128266, 7-token SNAC frames (pos0->L1, pos1/4->L2,
pos2/3/5/6->L3, offsets k*4096). Named-speaker prompt: [SOH] "voice: text" [EOT][SOA].
Zero-shot cloning: [BOS][SOH] ref_text [EOT][SOA][SOS] <ref audio tokens> [EOS_sp] then
[SOH] text [EOT][SOA]. Generation stops at end-of-speech (128258).
"""
import os, io, base64
import numpy as np, torch, soundfile as sf, requests
from scipy.signal import resample_poly
from fastapi import FastAPI, HTTPException
from fastapi.responses import Response
from pydantic import BaseModel
from transformers import AutoTokenizer
from snac import SNAC
MODEL_DIR = os.environ.get("MORPHEUS_MODEL_DIR", "/model")
SNAC_DIR = os.environ.get("SNAC_DIR", "/snac")
VLLM_URL = os.environ.get("VLLM_URL", "http://vllm-morpheus:8000/v1/completions")
VLLM_MODEL = os.environ.get("VLLM_MODEL", "morpheus")
SNAC_DEVICE = os.environ.get("SNAC_DEVICE", "cpu")
DEFAULT_VOICE = os.environ.get("MORPHEUS_DEFAULT_VOICE", "baddy")
VOICES = [v for v in os.environ.get("MORPHEUS_VOICES", "baddy").split(",") if v]
AUDIO_MAX = 156937
AUDIO_BASE, SOS, EOS_SP, SOH, SOA, EOT, BOS = 128266, 128257, 128258, 128259, 128260, 128009, 128000
tok = AutoTokenizer.from_pretrained(MODEL_DIR)
snac_model = SNAC.from_pretrained(SNAC_DIR).to(SNAC_DEVICE).eval()
app = FastAPI(title="mOrpheus TTS", version="0.1.0")
class TTSReq(BaseModel):
text: str
voice: str = DEFAULT_VOICE
temperature: float = 0.6
top_p: float = 0.95
max_tokens: int = 1200
repetition_penalty: float = 1.1 # keep <=1.1 for cloning (higher penalizes ref audio tokens)
reference_audio_b64: str | None = None # optional zero-shot clone: base64 WAV
reference_text: str | None = None # transcript of the reference
def _encode_ref(wav_bytes: bytes, ref_text: str) -> list[int]:
wav, sr = sf.read(io.BytesIO(wav_bytes))
if wav.ndim > 1:
wav = wav.mean(1)
wav = wav.astype(np.float32)
if sr != 24000:
wav = resample_poly(wav, 24000, sr).astype(np.float32)
wt = torch.tensor(wav, device=SNAC_DEVICE).view(1, 1, -1)
with torch.inference_mode():
codes = snac_model.encode(wt)
L1, L2, L3 = [c.squeeze(0).tolist() for c in codes]
ids = []
for i in range(len(L1)):
ids += [L1[i] + AUDIO_BASE, L2[2 * i] + AUDIO_BASE + 4096, L3[4 * i] + AUDIO_BASE + 8192,
L3[4 * i + 1] + AUDIO_BASE + 12288, L2[2 * i + 1] + AUDIO_BASE + 16384,
L3[4 * i + 2] + AUDIO_BASE + 20480, L3[4 * i + 3] + AUDIO_BASE + 24576]
return [BOS, SOH] + tok(ref_text, add_special_tokens=False).input_ids + [EOT, SOA, SOS] + ids + [EOS_SP]
def _build_prompt(req: TTSReq) -> list[int]:
if req.reference_audio_b64 and req.reference_text:
ref = _encode_ref(base64.b64decode(req.reference_audio_b64), req.reference_text)
return ref + [SOH] + tok(req.text, add_special_tokens=False).input_ids + [EOT, SOA]
return [SOH] + tok(f"{req.voice}: {req.text}").input_ids + [EOT, SOA]
def _decode(ids: list[int]):
if SOS in ids:
ids = ids[len(ids) - 1 - ids[::-1].index(SOS) + 1:]
codes = [t - AUDIO_BASE for t in ids if AUDIO_BASE <= t <= AUDIO_MAX]
l1, l2, l3 = [], [], []
for i in range(len(codes) // 7):
f = codes[7 * i:7 * i + 7]
c = [f[0], f[1] - 4096, f[2] - 8192, f[3] - 12288, f[4] - 16384, f[5] - 20480, f[6] - 24576]
if any(x < 0 or x > 4095 for x in c):
continue
l1.append(c[0]); l2 += [c[1], c[4]]; l3 += [c[2], c[3], c[5], c[6]]
if not l1:
return None
ct = [torch.tensor(l1).unsqueeze(0).to(SNAC_DEVICE),
torch.tensor(l2).unsqueeze(0).to(SNAC_DEVICE),
torch.tensor(l3).unsqueeze(0).to(SNAC_DEVICE)]
with torch.inference_mode():
return snac_model.decode(ct).squeeze().cpu().numpy()
@app.get("/health")
def health():
return {"status": "ok", "voices": VOICES, "default": DEFAULT_VOICE, "engine": VLLM_URL}
@app.get("/voices")
def voices():
return {"voices": VOICES, "default": DEFAULT_VOICE}
@app.post("/tts")
def tts(req: TTSReq):
prompt_ids = _build_prompt(req)
payload = {"model": VLLM_MODEL, "prompt": prompt_ids, "max_tokens": req.max_tokens,
"temperature": req.temperature, "top_p": req.top_p, "skip_special_tokens": False,
"stop_token_ids": [EOS_SP], "repetition_penalty": req.repetition_penalty}
try:
r = requests.post(VLLM_URL, json=payload, timeout=180)
except requests.RequestException as e:
raise HTTPException(502, f"engine unreachable: {e}")
if r.status_code != 200:
raise HTTPException(502, f"engine {r.status_code}: {r.text[:200]}")
text = r.json()["choices"][0]["text"]
gen_ids = tok(text, add_special_tokens=False).input_ids
audio = _decode(gen_ids)
if audio is None:
raise HTTPException(500, "no audio tokens generated")
buf = io.BytesIO()
sf.write(buf, audio, 24000, format="WAV", subtype="PCM_16")
return Response(content=buf.getvalue(), media_type="audio/wav",
headers={"X-Audio-Seconds": f"{len(audio)/24000:.2f}"})
+8
View File
@@ -0,0 +1,8 @@
snac
transformers
soundfile
scipy
numpy
fastapi
uvicorn[standard]
requests