Synchronize repository with Athena deployment

This commit is contained in:
Mikei386
2026-09-13 20:01:36 +02:00
parent 040a2df48b
commit fce9900389
60 changed files with 2570 additions and 607 deletions
+3 -1
View File
@@ -13,9 +13,11 @@ RUN apt-get update && apt-get install -y --no-install-recommends python3.12-venv
"transformers==${TRANSFORMERS_VERSION}" \
"accelerate==${ACCELERATE_VERSION}" \
"huggingface-hub==${HF_HUB_VERSION}" \
"nvidia-modelopt==0.46.0" bitsandbytes \
sentencepiece protobuf safetensors pillow && \
useradd --system --uid 10002 --home /nonexistent --shell /usr/sbin/nologin image-worker
COPY image_worker.py /app/image_worker.py
COPY image_worker_9b.py /app/image_worker_9b.py
USER 10002:10002
ENTRYPOINT ["/opt/image-venv/bin/python", "/app/image_worker.py"]
ENTRYPOINT ["/opt/image-venv/bin/python", "/app/image_worker_9b.py"]
@@ -0,0 +1,282 @@
#!/usr/bin/env python3
"""Private FLUX.2 Klein 9B FP8 beta worker for Athena's two GPUs.
The FP8 diffusion transformer runs on the RTX 5080. A Qwen3-8B NF4 text
encoder runs on the RTX 3060 while the profile controller temporarily pauses
Qwen3-TTS. The transformer and encoder are released before VAE decoding so
the 1024px decoder has sufficient workspace on the RTX 5080.
"""
from __future__ import annotations
import gc
import json
import os
import signal
import time
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from types import MethodType
HOST = os.environ.get("WORKER_HOST", "0.0.0.0")
PORT = int(os.environ.get("WORKER_PORT", "8086"))
TOKEN = os.environ.get("WORKER_TOKEN", "").strip()
COMPONENT_DIR = os.environ.get("FLUX_COMPONENT_DIR", "/models/components")
TRANSFORMER_FILE = os.environ.get(
"FLUX_TRANSFORMER_FILE", "/models/fp8/flux-2-klein-9b-fp8.safetensors")
OUTPUT_DIR = Path(os.environ.get("IMAGE_DIR", "/data/images")).resolve()
ACTIVE = False
os.environ.setdefault("DIFFUSERS_VERBOSITY", "error")
os.environ.setdefault("TRANSFORMERS_VERBOSITY", "error")
os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
if len(TOKEN) < 32:
raise RuntimeError("WORKER_TOKEN is missing or too short")
signal.signal(signal.SIGTERM, lambda *_: os._exit(0))
def _devices(torch):
if torch.cuda.device_count() != 2:
raise RuntimeError("FLUX 9B beta requires exactly two visible CUDA GPUs")
totals = {i: torch.cuda.get_device_properties(i).total_memory
for i in range(torch.cuda.device_count())}
transformer_index = max(totals, key=totals.get)
encoder_index = min(totals, key=totals.get)
return (transformer_index, encoder_index,
torch.device(f"cuda:{transformer_index}"),
torch.device(f"cuda:{encoder_index}"))
def _install_fp8_converter():
import diffusers.loaders.single_file_model as single_file_model
original = single_file_model.SINGLE_FILE_LOADABLE_CLASSES[
"Flux2Transformer2DModel"]["checkpoint_mapping_fn"]
scales = {}
double_map = {
"img_attn.proj": "attn.to_out.0",
"img_mlp.0": "ff.linear_in",
"img_mlp.2": "ff.linear_out",
"txt_attn.proj": "attn.to_add_out",
"txt_mlp.0": "ff_context.linear_in",
"txt_mlp.2": "ff_context.linear_out",
}
single_map = {
"linear1": "attn.to_qkv_mlp_proj",
"linear2": "attn.to_out",
}
def record(key, value):
parts = key.split(".")
scale_name, block = parts[-1], parts[1]
within = ".".join(parts[2:-1])
if parts[0] == "double_blocks":
if within == "img_attn.qkv":
targets = ("attn.to_q", "attn.to_k", "attn.to_v")
elif within == "txt_attn.qkv":
targets = ("attn.add_q_proj", "attn.add_k_proj",
"attn.add_v_proj")
else:
targets = (double_map[within],)
prefix = f"transformer_blocks.{block}"
elif parts[0] == "single_blocks":
targets = (single_map[within],)
prefix = f"single_transformer_blocks.{block}"
else:
raise ValueError(f"unexpected FP8 scale key: {key}")
for target in targets:
scales.setdefault(f"{prefix}.{target}", {})[scale_name] = value.clone()
def convert(checkpoint, **kwargs):
scales.clear()
for key in list(checkpoint):
if key.endswith((".input_scale", ".weight_scale")):
record(key, checkpoint.pop(key))
return original(checkpoint=checkpoint, **kwargs)
single_file_model.SINGLE_FILE_LOADABLE_CLASSES[
"Flux2Transformer2DModel"]["checkpoint_mapping_fn"] = convert
return scales
def _fp8_forward(torch, module, inputs):
shape = inputs.shape
input_fp8 = ((inputs / module._fp8_input_scale)
.clamp(torch.finfo(torch.float8_e4m3fn).min,
torch.finfo(torch.float8_e4m3fn).max)
.to(torch.float8_e4m3fn).reshape(-1, shape[-1]))
output = torch._scaled_mm(
input_fp8,
module.weight.reshape(-1, module.weight.shape[-1]).t(),
scale_a=module._fp8_input_scale,
scale_b=module._fp8_weight_scale,
bias=module.bias,
out_dtype=inputs.dtype,
use_fast_accum=True,
)
return output.reshape(*shape[:-1], output.shape[-1])
def generate(data: dict) -> dict:
global ACTIVE
import torch
from diffusers import (Flux2KleinPipeline, Flux2Transformer2DModel,
NVIDIAModelOptConfig)
from modelopt.torch.opt import enable_huggingface_checkpointing
from modelopt.torch.quantization.config import FP8_DEFAULT_CFG
from PIL import Image
from transformers import BitsAndBytesConfig, Qwen3ForCausalLM
prompt, filename = data.get("prompt"), data.get("filename")
if not isinstance(prompt, str) or not prompt.strip() or len(prompt) > 8000:
raise ValueError("invalid prompt")
if (not isinstance(filename, str) or Path(filename).name != filename
or not filename.endswith(".png")):
raise ValueError("invalid filename")
width, height = int(data.get("width", 1024)), int(data.get("height", 1024))
if (width, height) != (1024, 1024):
raise ValueError("FLUX 9B beta currently supports only 1024x1024")
if int(data.get("steps", 4)) != 4 or float(data.get("guidance", 1.0)) != 1.0:
raise ValueError("FLUX 9B beta requires steps=4 and guidance=1.0")
source_files = data.get("source_files") or []
if not isinstance(source_files, list) or len(source_files) > 4:
raise ValueError("invalid source image list")
source_images = []
for name in source_files:
if not isinstance(name, str) or Path(name).name != name:
raise ValueError("invalid source image filename")
source = (OUTPUT_DIR / name).resolve()
if source.parent != OUTPUT_DIR or not source.is_file():
raise ValueError("source image not found")
with Image.open(source) as opened:
source_images.append(opened.convert("RGB"))
started = time.monotonic()
ACTIVE = True
transformer = text_encoder = pipe = latent = decoded = image = None
try:
enable_huggingface_checkpointing()
scales = _install_fp8_converter()
tx_index, enc_index, tx_device, enc_device = _devices(torch)
quantization = NVIDIAModelOptConfig(
quant_type="FP8", weight_only=False,
modelopt_config=FP8_DEFAULT_CFG)
transformer = Flux2Transformer2DModel.from_single_file(
TRANSFORMER_FILE, config=COMPONENT_DIR, subfolder="transformer",
quantization_config=quantization, torch_dtype=torch.bfloat16,
device_map={"": tx_index}, local_files_only=True)
patched = 0
for module_name, module in transformer.named_modules():
if module_name not in scales:
continue
module.register_buffer("_fp8_input_scale",
scales[module_name]["input_scale"])
module.register_buffer("_fp8_weight_scale",
scales[module_name]["weight_scale"])
module.forward = MethodType(
lambda self, inputs: _fp8_forward(torch, self, inputs), module)
patched += 1
if patched != len(scales):
raise RuntimeError(f"patched only {patched} of {len(scales)} FP8 layers")
transformer.to(tx_device)
text_encoder = Qwen3ForCausalLM.from_pretrained(
os.path.join(COMPONENT_DIR, "text_encoder"),
torch_dtype=torch.bfloat16, low_cpu_mem_usage=True,
quantization_config=BitsAndBytesConfig(
load_in_4bit=True, bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16,
bnb_4bit_use_double_quant=True),
device_map={"": enc_index}, local_files_only=True)
pipe = Flux2KleinPipeline.from_pretrained(
COMPONENT_DIR, transformer=transformer, text_encoder=text_encoder,
torch_dtype=torch.bfloat16, local_files_only=True)
pipe.vae.enable_slicing()
pipe.vae.enable_tiling()
pipe.vae.to(tx_device)
loaded = time.monotonic() - started
prompt_embeds, _ = pipe.encode_prompt(
prompt.strip(), device=enc_device, max_sequence_length=128)
prompt_embeds = prompt_embeds.to(tx_device)
pipe.text_encoder = None
seed = data.get("seed")
generator = None if seed is None else torch.Generator(
device=tx_device).manual_seed(int(seed))
kwargs = {
"prompt": None, "prompt_embeds": prompt_embeds,
"height": height, "width": width, "num_inference_steps": 4,
"guidance_scale": 1.0, "generator": generator,
"output_type": "latent",
}
if source_images:
kwargs["image"] = (source_images[0] if len(source_images) == 1
else source_images)
latent = pipe(**kwargs).images
pipe.transformer = None
del transformer, text_encoder, prompt_embeds, generator
transformer = text_encoder = None
gc.collect()
torch.cuda.empty_cache()
latent = latent.to(device=tx_device, dtype=pipe.vae.dtype)
decoded = pipe.vae.decode(latent, return_dict=False)[0]
image = pipe.image_processor.postprocess(
decoded.detach(), output_type="pil")[0]
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
image.save(OUTPUT_DIR / filename)
return {"status": "ok", "filename": filename,
"seconds": round(time.monotonic() - started, 3),
"load_seconds": round(loaded, 3),
"model": "FLUX.2-klein-9B-fp8-beta"}
finally:
for value in (image, decoded, latent, pipe, text_encoder, transformer):
if value is not None:
del value
gc.collect()
torch.cuda.empty_cache()
ACTIVE = False
class Handler(BaseHTTPRequestHandler):
def log_message(self, fmt: str, *args: object) -> None:
print(f"[flux9b-beta] {self.client_address[0]} {fmt % args}", flush=True)
def reply(self, status: int, payload: dict) -> None:
body = json.dumps(payload, separators=(",", ":")).encode()
self.send_response(status)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def do_GET(self) -> None: # noqa: N802
if self.path == "/health":
self.reply(200, {"status": "ok", "model_loaded": ACTIVE,
"model": "FLUX.2-klein-9B-fp8-beta"})
else:
self.reply(404, {"error": "not found"})
def do_POST(self) -> None: # noqa: N802
if self.headers.get("Authorization", "") != f"Bearer {TOKEN}":
self.reply(401, {"error": "unauthorized"})
return
if self.path != "/generate":
self.reply(404, {"error": "not found"})
return
try:
length = int(self.headers.get("Content-Length", "0"))
if length < 2 or length > 16384:
raise ValueError("invalid request size")
self.reply(200, generate(json.loads(self.rfile.read(length))))
except Exception as exc:
print(f"[flux9b-beta] generation failed: {type(exc).__name__}: "
f"{str(exc)[:1000]}", flush=True)
self.reply(500, {"status": "error", "message": str(exc)})
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
-24
View File
@@ -1,24 +0,0 @@
FROM python:3.12-slim-bookworm
ARG PIPER_TTS_VERSION=1.6.0
RUN apt-get update \
&& apt-get install -y --no-install-recommends ca-certificates curl ffmpeg gosu \
&& python -m pip install --no-cache-dir "piper-tts==${PIPER_TTS_VERSION}" \
&& useradd --system --uid 10003 --home-dir /nonexistent --shell /usr/sbin/nologin piper \
&& rm -rf /var/lib/apt/lists/*
WORKDIR /app
COPY piper_worker.py /app/piper_worker.py
COPY entrypoint.sh /usr/local/bin/mike-ai-piper-entrypoint
RUN chmod 0755 /usr/local/bin/mike-ai-piper-entrypoint
ENV PIPER_DATA_DIR=/data \
PIPER_VOICE=de_DE-thorsten-high \
PIPER_VOICE_ALIAS=alloy \
PIPER_HOST=0.0.0.0 \
PIPER_PORT=8085
VOLUME ["/data"]
EXPOSE 8085
ENTRYPOINT ["/usr/local/bin/mike-ai-piper-entrypoint"]
-15
View File
@@ -1,15 +0,0 @@
#!/bin/sh
set -eu
data_dir=${PIPER_DATA_DIR:-/data}
voice=${PIPER_VOICE:-de_DE-thorsten-high}
mkdir -p "$data_dir"
chown 10003:10003 "$data_dir"
if [ ! -s "$data_dir/$voice.onnx" ] || [ ! -s "$data_dir/$voice.onnx.json" ]; then
echo "Downloading Piper voice: $voice"
gosu piper python -m piper.download_voices --data-dir "$data_dir" "$voice"
fi
exec gosu piper python /app/piper_worker.py
-153
View File
@@ -1,153 +0,0 @@
#!/usr/bin/env python3
"""Small, private Piper worker for the Mike AI profile router.
The public OpenAI-compatible endpoint remains in the router. This worker only
accepts the narrow internal /status and /tts protocol and never logs input text.
"""
from __future__ import annotations
import io
import json
import os
import subprocess
import threading
import wave
from http import HTTPStatus
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from piper import PiperVoice, SynthesisConfig
DATA_DIR = Path(os.getenv("PIPER_DATA_DIR", "/data"))
VOICE_NAME = os.getenv("PIPER_VOICE", "de_DE-thorsten-high")
VOICE_ALIAS = os.getenv("PIPER_VOICE_ALIAS", "alloy")
HOST = os.getenv("PIPER_HOST", "0.0.0.0")
PORT = int(os.getenv("PIPER_PORT", "8085"))
MAX_TEXT_CHARS = int(os.getenv("PIPER_MAX_TEXT_CHARS", "8000"))
MAX_REQUEST_BYTES = int(os.getenv("PIPER_MAX_REQUEST_BYTES", "65536"))
VOICE_PATH = DATA_DIR / f"{VOICE_NAME}.onnx"
VOICE = PiperVoice.load(str(VOICE_PATH))
SYNTHESIS_LOCK = threading.Lock()
def synthesize_wav(text: str, speed: float) -> bytes:
"""Synthesize a complete WAV in memory without retaining the text."""
output = io.BytesIO()
config = SynthesisConfig(length_scale=1.0 / speed)
with SYNTHESIS_LOCK, wave.open(output, "wb") as wav_file:
VOICE.synthesize_wav(text, wav_file, syn_config=config)
return output.getvalue()
def wav_to_mp3(wav_bytes: bytes) -> bytes:
"""Convert Piper's WAV to the MP3 format Open WebUI requests by default."""
result = subprocess.run(
[
"ffmpeg", "-hide_banner", "-loglevel", "error",
"-f", "wav", "-i", "pipe:0",
"-codec:a", "libmp3lame", "-b:a", "96k",
"-f", "mp3", "pipe:1",
],
input=wav_bytes,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
check=False,
timeout=120,
)
if result.returncode != 0:
raise RuntimeError("ffmpeg conversion failed")
return result.stdout
class Handler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def log_message(self, fmt: str, *args: object) -> None:
# Deliberately omit URLs and request bodies from the log.
print(f"piper-worker: {self.command} -> {args[1] if len(args) > 1 else '-'}")
def send_bytes(self, status: int, body: bytes, content_type: str) -> None:
self.send_response(status)
self.send_header("Content-Type", content_type)
self.send_header("Content-Length", str(len(body)))
self.send_header("Cache-Control", "no-store")
self.end_headers()
self.wfile.write(body)
def send_json(self, status: int, payload: dict) -> None:
self.send_bytes(
status,
json.dumps(payload, separators=(",", ":")).encode(),
"application/json",
)
def do_GET(self) -> None: # noqa: N802
if self.path != "/status":
self.send_json(HTTPStatus.NOT_FOUND, {"error": "not found"})
return
self.send_json(
HTTPStatus.OK,
{
"ready": True,
"engine": "piper",
"model": VOICE_NAME,
"voices": [VOICE_ALIAS],
},
)
def do_POST(self) -> None: # noqa: N802
if self.path != "/tts":
self.send_json(HTTPStatus.NOT_FOUND, {"error": "not found"})
return
try:
content_length = int(self.headers.get("Content-Length", "0"))
except ValueError:
content_length = 0
if content_length <= 0 or content_length > MAX_REQUEST_BYTES:
self.send_json(HTTPStatus.REQUEST_ENTITY_TOO_LARGE, {"error": "invalid request size"})
return
try:
request = json.loads(self.rfile.read(content_length))
text = request.get("text", "")
voice = request.get("voice", VOICE_ALIAS)
output_format = request.get("format", "mp3")
speed = float(request.get("speed", 1.0))
except (json.JSONDecodeError, TypeError, ValueError):
self.send_json(HTTPStatus.BAD_REQUEST, {"error": "invalid JSON request"})
return
if not isinstance(text, str) or not text.strip() or len(text) > MAX_TEXT_CHARS:
self.send_json(HTTPStatus.BAD_REQUEST, {"error": "invalid text"})
return
if voice != VOICE_ALIAS:
self.send_json(HTTPStatus.BAD_REQUEST, {"error": "unknown voice"})
return
if output_format not in {"wav", "mp3"}:
self.send_json(HTTPStatus.BAD_REQUEST, {"error": "unsupported format"})
return
if not 0.5 <= speed <= 2.0:
self.send_json(HTTPStatus.BAD_REQUEST, {"error": "invalid speed"})
return
try:
audio = synthesize_wav(text.strip(), speed)
if output_format == "mp3":
audio = wav_to_mp3(audio)
content_type = "audio/mpeg"
else:
content_type = "audio/wav"
except (OSError, RuntimeError, subprocess.SubprocessError):
self.send_json(HTTPStatus.INTERNAL_SERVER_ERROR, {"error": "synthesis failed"})
return
self.send_bytes(HTTPStatus.OK, audio, content_type)
if __name__ == "__main__":
print(f"Piper worker ready: {VOICE_NAME} as {VOICE_ALIAS} on {HOST}:{PORT}")
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
+27 -13
View File
@@ -132,6 +132,25 @@ class LanguageSegmentationTests(unittest.TestCase):
self.assertIn("5 bis 6 Uhr", spoken)
self.assertIn("Wind maximal 16 Kilometer pro Stunde", spoken)
def test_qwen_speaks_aspect_ratios_as_ratios(self):
spoken = gateway.prepare_for_qwen_speech(
"Cover sind im Hochformat (2:3), Screenshots im Querformat "
"(16:9), ein Quadrat im Seitenverhältnis 2:2 und 4:3-Format."
)
self.assertIn("Hochformat (2 zu 3)", spoken)
self.assertIn("Querformat (16 zu 9)", spoken)
self.assertIn("Seitenverhältnis 2 zu 2", spoken)
self.assertIn("4 zu 3-Format", spoken)
def test_qwen_keeps_clock_times_distinct_from_aspect_ratios(self):
spoken = gateway.prepare_for_qwen_speech(
"Beginn um 16:09 Uhr, Fehler um 02:14; das Videoformat ist 16:9."
)
self.assertIn("16 Uhr 9", spoken)
self.assertNotIn("16 Uhr 9 Uhr", spoken)
self.assertIn("2 Uhr 14", spoken)
self.assertIn("Videoformat ist 16 zu 9", spoken)
def test_qwen_speaks_strict_date_ranges_as_calendar_dates(self):
spoken = gateway.prepare_for_qwen_speech(
"Neuigkeiten vom 04.–05.09. und Vergleich 04.09.–06.10.2026."
@@ -225,25 +244,20 @@ class LanguageSegmentationTests(unittest.TestCase):
self.assertTrue(all(len(part.split()) > 4 for _, part in segments))
class FallbackTests(unittest.TestCase):
class BackendFailureTests(unittest.TestCase):
def setUp(self):
self.original_xtts = gateway.synthesize_xtts
self.original_piper = gateway.synthesize_piper
self.original_qwen = gateway.synthesize_qwen
def tearDown(self):
gateway.synthesize_xtts = self.original_xtts
gateway.synthesize_piper = self.original_piper
gateway.synthesize_qwen = self.original_qwen
def test_piper_is_used_when_xtts_fails(self):
def test_qwen_failure_is_reported_without_fallback(self):
def fail(*_args):
raise RuntimeError("synthetic XTTS failure")
raise RuntimeError("synthetic Qwen failure")
gateway.synthesize_xtts = fail
gateway.synthesize_piper = lambda *_args: (b"piper", "audio/wav")
self.assertEqual(
gateway.synthesize("synthetic test", "wav", 1.0),
(b"piper", "audio/wav"),
)
gateway.synthesize_qwen = fail
with self.assertRaisesRegex(RuntimeError, "synthetic Qwen failure"):
gateway.synthesize("synthetic test", "wav", 1.0)
class AudioJoinTests(unittest.TestCase):
+166 -36
View File
@@ -1,5 +1,5 @@
#!/usr/bin/env python3
"""Private Qwen3-TTS-first gateway with a Piper fallback.
"""Private Qwen3-TTS gateway.
The gateway implements the narrow /status and /tts protocol already consumed
by the profile router. Request text is never logged or persisted.
@@ -8,6 +8,7 @@ by the profile router. Request text is never logged or persisted.
from __future__ import annotations
import io
import http.client
import json
import os
import re
@@ -17,6 +18,7 @@ import time
import unicodedata
import urllib.error
import urllib.request
import urllib.parse
import wave
from array import array
from http import HTTPStatus
@@ -31,7 +33,6 @@ QWEN_TTS_VOICE = os.getenv("QWEN_TTS_VOICE", "serena")
QWEN_TTS_LANGUAGE = os.getenv("QWEN_TTS_LANGUAGE", "German")
QWEN_TTS_TIMEOUT = float(os.getenv("QWEN_TTS_TIMEOUT", "120"))
XTTS_URL = os.getenv("XTTS_URL", "http://xtts:80").rstrip("/")
PIPER_URL = os.getenv("PIPER_URL", "http://piper:8085").rstrip("/")
VOICE_ALIAS = os.getenv("TTS_VOICE_ALIAS", "alloy")
XTTS_SPEAKER = os.getenv("XTTS_SPEAKER", "Annmarie Nele")
DEFAULT_LANGUAGE = os.getenv("TTS_DEFAULT_LANGUAGE", "de")
@@ -41,7 +42,6 @@ MAX_TEXT_CHARS = int(os.getenv("TTS_MAX_TEXT_CHARS", "8000"))
MAX_REQUEST_BYTES = int(os.getenv("TTS_MAX_REQUEST_BYTES", "65536"))
MAX_AUDIO_BYTES = int(os.getenv("TTS_MAX_AUDIO_BYTES", str(64 * 1024 * 1024)))
XTTS_TIMEOUT = float(os.getenv("XTTS_TIMEOUT", "120"))
PIPER_TIMEOUT = float(os.getenv("PIPER_TIMEOUT", "120"))
QUEUE_TIMEOUT = float(os.getenv("XTTS_QUEUE_TIMEOUT", "15"))
# XTTS loses natural prosody when a sentence is synthesized as many tiny
# requests: every request starts a fresh utterance. Keep complete sentences
@@ -61,8 +61,7 @@ SPEAKER_LOCK = threading.Lock()
SPEAKER_CONDITIONING: dict | None = None
STATE = {
"last_backend": None,
"xtts_failures": 0,
"piper_fallbacks": 0,
"qwen_failures": 0,
"last_error": None,
}
@@ -267,6 +266,36 @@ def _spoken_ipv4(match: re.Match) -> str:
return " Punkt ".join(str(int(part)) for part in match.group(0).split("."))
def _normalize_aspect_ratios(text: str) -> str:
"""Speak colon notation as a ratio only when the surrounding text says so.
A bare ``16:09`` remains a clock time. This deliberately avoids a global
replacement of common ratios because ``16:9`` can also be a valid time.
"""
cue = (
r"(?:Seitenverh[aä]ltnis|Bildseitenverh[aä]ltnis|Bildformat|"
r"Videoformat|Hochformat|Querformat|Format|Aspect[- ]?Ratio)"
)
text = re.sub(
rf"\b({cue}\b(?:\s+(?:von|im|ist|betr[aä]gt))?\s*[\(\[]?\s*)"
rf"(\d{{1,3}})\s*:\s*(\d{{1,3}})",
lambda match: (
f"{match.group(1)}{int(match.group(2))} zu {int(match.group(3))}"
),
text,
flags=re.IGNORECASE,
)
return re.sub(
rf"\b(\d{{1,3}})\s*:\s*(\d{{1,3}})"
rf"(\s*[-‐‑‒–—−]?\s*{cue}\b)",
lambda match: (
f"{int(match.group(1))} zu {int(match.group(2))}{match.group(3)}"
),
text,
flags=re.IGNORECASE,
)
def normalize_for_german_speech(text: str) -> str:
"""Turn common visual notation into unambiguous spoken German."""
text = re.sub(r"\bv\.\s*a\.", "vor allem", text, flags=re.IGNORECASE)
@@ -279,10 +308,14 @@ def normalize_for_german_speech(text: str) -> str:
_spoken_ipv4,
text,
)
# A colon is ambiguous between an aspect ratio and a clock time. Resolve
# ratios first, but only when an explicit format cue is present.
text = _normalize_aspect_ratios(text)
text = re.sub(
r"\b([01]?\d|2[0-3]):([0-5]\d)\b",
r"\b([01]?\d|2[0-3]):([0-5]\d)\b(?:\s*Uhr\b)?",
lambda match: f"{int(match.group(1))} Uhr {int(match.group(2))}",
text,
flags=re.IGNORECASE,
)
text = re.sub(
r"\b([01]?\d|2[0-3])\s*[-‐‑‒–—−]\s*"
@@ -702,18 +735,6 @@ def synthesize_xtts(text: str, output_format: str,
return _convert(_wav(_join_pcm(pcm_parts)), output_format, speed)
def synthesize_piper(text: str, output_format: str,
speed: float) -> tuple[bytes, str]:
upstream_format = "wav" if output_format == "pcm" else output_format
audio, content_type = _request(
f"{PIPER_URL}/tts",
payload={"text": text, "voice": "alloy", "speed": speed,
"format": upstream_format},
timeout=PIPER_TIMEOUT,
)
return _convert(audio, "pcm", 1.0) if output_format == "pcm" else (audio, content_type)
def synthesize_qwen(text: str, output_format: str,
speed: float) -> tuple[bytes, str]:
text = prepare_for_qwen_speech(text)
@@ -728,6 +749,47 @@ def synthesize_qwen(text: str, output_format: str,
return _convert(audio, "pcm", 1.0) if output_format == "pcm" else (audio, content_type)
def open_qwen_pcm_stream(text: str, chunk_size: int = 4) \
-> tuple[http.client.HTTPConnection, http.client.HTTPResponse]:
"""Open Qwen's native token-level PCM stream without buffering it.
The upstream emits headerless 24 kHz mono signed 16-bit little-endian
PCM. Keeping this response streaming is what lets playback begin while
the remainder of the sentence is still being synthesized.
"""
parsed = urllib.parse.urlparse(QWEN_TTS_URL)
if parsed.scheme != "http" or not parsed.hostname:
raise RuntimeError("QWEN_TTS_URL must be an http URL")
port = parsed.port or 80
prefix = parsed.path.rstrip("/")
payload = json.dumps({
"model": QWEN_TTS_MODEL,
"input": prepare_for_qwen_speech(text),
"voice": QWEN_TTS_VOICE,
"language": QWEN_TTS_LANGUAGE,
"chunk_size": chunk_size,
}, separators=(",", ":")).encode()
connection = http.client.HTTPConnection(
parsed.hostname, port, timeout=QWEN_TTS_TIMEOUT)
try:
connection.request(
"POST",
f"{prefix}/v1/audio/speech/pcm-stream",
body=payload,
headers={"Content-Type": "application/json",
"Accept": "application/octet-stream"},
)
response = connection.getresponse()
if response.status != HTTPStatus.OK:
message = response.read(512).decode(errors="replace")
raise RuntimeError(
f"Qwen PCM stream failed ({response.status}): {message}")
return connection, response
except Exception:
connection.close()
raise
def synthesize(text: str, output_format: str, speed: float) -> tuple[bytes, str]:
acquired = SYNTHESIS_LOCK.acquire(timeout=QUEUE_TIMEOUT)
if acquired:
@@ -737,22 +799,18 @@ def synthesize(text: str, output_format: str, speed: float) -> tuple[bytes, str]
STATE["last_backend"] = "qwen3-tts-1.7b"
STATE["last_error"] = None
return audio
except Exception as exc: # fallback must cover all Qwen failures
except Exception as exc:
with STATE_LOCK:
STATE["xtts_failures"] += 1
STATE["qwen_failures"] += 1
STATE["last_error"] = type(exc).__name__
raise
finally:
SYNTHESIS_LOCK.release()
else:
with STATE_LOCK:
STATE["xtts_failures"] += 1
STATE["qwen_failures"] += 1
STATE["last_error"] = "queue-timeout"
audio = synthesize_piper(text, output_format, speed)
with STATE_LOCK:
STATE["last_backend"] = "piper"
STATE["piper_fallbacks"] += 1
return audio
raise RuntimeError("speech queue timeout")
class Handler(BaseHTTPRequestHandler):
@@ -779,25 +837,26 @@ class Handler(BaseHTTPRequestHandler):
self.send_json(HTTPStatus.NOT_FOUND, {"error": "not found"})
return
primary_ready = _reachable(QWEN_TTS_URL, "/health")
fallback_ready = _reachable(PIPER_URL, "/status")
with STATE_LOCK:
state = dict(STATE)
# This endpoint is also the container liveness check. Qwen3-TTS is
# deliberately stopped in exclusive GPU modes such as Applio, so the
# gateway itself must stay healthy while reporting ready=false.
self.send_json(
HTTPStatus.OK if fallback_ready else HTTPStatus.SERVICE_UNAVAILABLE,
HTTPStatus.OK,
{
"ready": fallback_ready,
"engine": "qwen3-tts-with-piper-fallback",
"ready": primary_ready,
"engine": "qwen3-tts",
"model": "Qwen3-TTS-12Hz-1.7B-Base",
"voices": [VOICE_ALIAS],
"speaker": QWEN_TTS_VOICE,
"primary_ready": primary_ready,
"fallback_ready": fallback_ready,
**state,
},
)
def do_POST(self) -> None: # noqa: N802
if self.path != "/tts":
if self.path not in {"/tts", "/tts/pcm-stream"}:
self.send_json(HTTPStatus.NOT_FOUND, {"error": "not found"})
return
try:
@@ -810,7 +869,7 @@ class Handler(BaseHTTPRequestHandler):
return
try:
request = json.loads(self.rfile.read(length))
text = request.get("text", "")
text = request.get("input", request.get("text", ""))
voice = request.get("voice", VOICE_ALIAS)
output_format = request.get("format", "mp3")
speed = float(request.get("speed", 1.0))
@@ -829,6 +888,9 @@ class Handler(BaseHTTPRequestHandler):
if not 0.5 <= speed <= 2.0:
self.send_json(HTTPStatus.BAD_REQUEST, {"error": "invalid speed"})
return
if self.path == "/tts/pcm-stream":
self._stream_qwen_pcm(text.strip(), request)
return
started = time.monotonic()
try:
audio, content_type = synthesize(text.strip(), output_format, speed)
@@ -836,13 +898,81 @@ class Handler(BaseHTTPRequestHandler):
with STATE_LOCK:
STATE["last_error"] = type(exc).__name__
self.send_json(HTTPStatus.SERVICE_UNAVAILABLE,
{"error": "all local speech backends failed"})
{"error": "local Qwen3-TTS backend failed"})
return
print(f"tts-gateway: synthesized via {STATE['last_backend']} in "
f"{time.monotonic() - started:.2f}s")
self.send_bytes(HTTPStatus.OK, audio, content_type)
def _stream_qwen_pcm(self, text: str, request: dict) -> None:
"""Unframe Qwen's PCM frames and relay their audio immediately."""
try:
chunk_size = max(1, min(32, int(request.get("chunk_size", 4))))
except (TypeError, ValueError):
self.send_json(HTTPStatus.BAD_REQUEST,
{"error": "invalid chunk_size"})
return
acquired = SYNTHESIS_LOCK.acquire(timeout=QUEUE_TIMEOUT)
if not acquired:
self.send_json(HTTPStatus.SERVICE_UNAVAILABLE,
{"error": "speech queue timeout"})
return
connection = None
started = time.monotonic()
headers_sent = False
try:
connection, response = open_qwen_pcm_stream(text, chunk_size)
self.send_response(HTTPStatus.OK)
self.send_header("Content-Type", "application/octet-stream")
self.send_header("Cache-Control", "no-store")
self.send_header("Connection", "close")
self.end_headers()
headers_sent = True
first = True
while True:
frame_header = response.read(4)
if not frame_header:
break
if len(frame_header) != 4:
raise RuntimeError("truncated Qwen PCM frame header")
frame_length = int.from_bytes(frame_header, "big")
if frame_length == 0:
break
if frame_length > MAX_AUDIO_BYTES:
raise RuntimeError("Qwen PCM frame is too large")
remaining = frame_length
while remaining:
chunk = response.read(min(16384, remaining))
if not chunk:
raise RuntimeError("truncated Qwen PCM frame")
if first:
print("tts-gateway: first Qwen PCM chunk in "
f"{time.monotonic() - started:.2f}s")
first = False
self.wfile.write(chunk)
self.wfile.flush()
remaining -= len(chunk)
with STATE_LOCK:
STATE["last_backend"] = "qwen3-tts-1.7b-stream"
STATE["last_error"] = None
except Exception as exc:
with STATE_LOCK:
STATE["last_error"] = type(exc).__name__
# Once PCM started, simply close the truncated response. Sending
# JSON into the audio stream would produce loud corrupt samples.
if not headers_sent and not self.wfile.closed:
try:
self.send_json(HTTPStatus.SERVICE_UNAVAILABLE,
{"error": "local PCM stream failed"})
except (OSError, BrokenPipeError):
pass
finally:
if connection is not None:
connection.close()
SYNTHESIS_LOCK.release()
self.close_connection = True
if __name__ == "__main__":
print(f"TTS gateway ready on {HOST}:{PORT}; primary={QWEN_TTS_VOICE}; fallback=Piper")
print(f"TTS gateway ready on {HOST}:{PORT}; backend=Qwen3-TTS; voice={QWEN_TTS_VOICE}")
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()