- Neues xtts_worker.py: Coqui XTTS-v2, HTTP-API auf Port 8085 - Router: TTS_WORKER_URL → 8085, TTS_MODEL → xtts-v2, TTS_VOICES → claribel - deploy: mike-ai-xtts.service, install.sh + deploy.sh aktualisiert - Tests: 54/54 bestanden (mock_tts_worker + test_local.sh auf XTTS umgestellt) - README: TTS-Section auf XTTS-v2 aktualisiert - Kokoro-Service gestoppt und deaktiviert (Dateien bleiben als Backup)
284 lines
8.5 KiB
Python
284 lines
8.5 KiB
Python
#!/usr/bin/env python3
|
||
"""XTTS-v2 TTS-Worker: langlebiger HTTP-Server, hält das Modell im RAM.
|
||
|
||
Endpunkte:
|
||
GET /status → Health-Check
|
||
POST /tts → Synthese (JSON: text, voice, speed, format)
|
||
|
||
CPU-only, keine GPU. Logging nach stdout (journald).
|
||
"""
|
||
|
||
import base64
|
||
import io
|
||
import json
|
||
import logging
|
||
import os
|
||
import sys
|
||
import time
|
||
import wave
|
||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||
|
||
# CPU-only erzwingen (BEVOR torch importiert wird)
|
||
os.environ["CUDA_VISIBLE_DEVICES"] = ""
|
||
|
||
import numpy as np
|
||
import torch
|
||
|
||
# XTTS-v2 Modell-ID
|
||
MODEL_ID = "tts_models/multilingual/multi-dataset/xtts_v2"
|
||
|
||
# Verfügbare Stimmen (XTTS Speaker-Namen)
|
||
VOICES = {
|
||
"claribel": "Claribel Dervla",
|
||
}
|
||
|
||
# Default-Stimme
|
||
DEFAULT_VOICE = "claribel"
|
||
|
||
# Port
|
||
PORT = int(os.environ.get("XTTS_PORT", "8085"))
|
||
|
||
# Logging nach stdout/stderr (für journald)
|
||
logging.basicConfig(
|
||
level=logging.INFO,
|
||
format="%(asctime)s %(levelname)s %(message)s",
|
||
stream=sys.stdout,
|
||
)
|
||
log = logging.getLogger("xtts-worker")
|
||
|
||
|
||
class XTTSWorker:
|
||
"""Hält das XTTS-v2-Modell geladen und synthetisiert Audio."""
|
||
|
||
def __init__(self):
|
||
self.tts = None
|
||
self.status = "loading"
|
||
self.error = None
|
||
self.model_name = "xtts-v2"
|
||
self.voices = list(VOICES.keys())
|
||
|
||
def load(self):
|
||
"""Lädt das XTTS-v2-Modell (einmalig)."""
|
||
log.info("Lade XTTS-v2-Modell: %s", MODEL_ID)
|
||
start = time.time()
|
||
try:
|
||
from TTS.api import TTS
|
||
|
||
# Modell laden (CPU-only)
|
||
self.tts = TTS(MODEL_ID)
|
||
self.status = "ready"
|
||
elapsed = time.time() - start
|
||
log.info("Modell geladen in %.1fs", elapsed)
|
||
except Exception as e:
|
||
self.status = "error"
|
||
self.error = str(e)
|
||
log.error("Fehler beim Laden: %s", e)
|
||
raise
|
||
|
||
def synthesize(
|
||
self,
|
||
text: str,
|
||
voice: str = DEFAULT_VOICE,
|
||
language: str = "de",
|
||
speed: float = 1.0,
|
||
output_format: str = "mp3",
|
||
) -> tuple[bytes, str]:
|
||
"""Synthetisiert Audio und gibt (bytes, content_type) zurück."""
|
||
if self.tts is None:
|
||
raise RuntimeError("Modell nicht geladen")
|
||
|
||
# Voice-Name auflösen
|
||
speaker_name = VOICES.get(voice, voice)
|
||
|
||
# WAV synthetisieren
|
||
start = time.time()
|
||
wav_data = self.tts.tts(
|
||
text=text,
|
||
speaker=speaker_name,
|
||
language=language,
|
||
)
|
||
synth_time = time.time() - start
|
||
|
||
# tts.tts() gibt eine Liste von Floats zurück (Audio-Samples)
|
||
audio_np = np.array(wav_data, dtype=np.float32)
|
||
if audio_np.ndim > 1:
|
||
audio_np = audio_np.squeeze()
|
||
|
||
# Normalisieren und zu int16 konvertieren
|
||
audio_np = audio_np / max(1e-8, np.abs(audio_np).max())
|
||
audio_np = (audio_np * 32767).astype(np.int16)
|
||
|
||
# WAV schreiben
|
||
sample_rate = 24000 # XTTS-v2 Sample Rate
|
||
wav_buffer = io.BytesIO()
|
||
with wave.open(wav_buffer, "wb") as wav_file:
|
||
wav_file.setnchannels(1)
|
||
wav_file.setsampwidth(2) # 16-bit
|
||
wav_file.setframerate(sample_rate)
|
||
wav_file.writeframes(audio_np.tobytes())
|
||
|
||
wav_bytes = wav_buffer.getvalue()
|
||
duration = len(audio_np) / sample_rate
|
||
|
||
# Speed anwenden (resample)
|
||
if speed != 1.0:
|
||
try:
|
||
import torchaudio
|
||
|
||
audio_tensor = torch.from_numpy(audio_np).float().unsqueeze(0)
|
||
resampler = torchaudio.transforms.Resample(
|
||
orig_freq=sample_rate,
|
||
new_freq=int(sample_rate * speed),
|
||
)
|
||
audio_tensor = resampler(audio_tensor)
|
||
audio_np = audio_tensor.squeeze(0).numpy()
|
||
audio_np = (audio_np * 32767).astype(np.int16)
|
||
|
||
wav_buffer = io.BytesIO()
|
||
with wave.open(wav_buffer, "wb") as wav_file:
|
||
wav_file.setnchannels(1)
|
||
wav_file.setsampwidth(2)
|
||
wav_file.setframerate(int(sample_rate * speed))
|
||
wav_file.writeframes(audio_np.tobytes())
|
||
wav_bytes = wav_buffer.getvalue()
|
||
duration = len(audio_np) / (sample_rate * speed)
|
||
except Exception as e:
|
||
log.warning("Speed-Resampling fehlgeschlagen: %s", e)
|
||
|
||
# Format konvertieren
|
||
if output_format == "mp3":
|
||
try:
|
||
import torchaudio
|
||
|
||
# WAV zu MP3 – audio_np ist int16, muss zu float32 Tensor
|
||
audio_float = torch.from_numpy(audio_np.astype(np.float32) / 32768.0).unsqueeze(0)
|
||
mp3_buffer = io.BytesIO()
|
||
torchaudio.save(
|
||
mp3_buffer,
|
||
audio_float,
|
||
sample_rate=int(sample_rate * speed) if speed != 1.0 else sample_rate,
|
||
format="mp3",
|
||
)
|
||
audio_bytes = mp3_buffer.getvalue()
|
||
content_type = "audio/mpeg"
|
||
except Exception as e:
|
||
log.warning("MP3-Konvertierung fehlgeschlagen, liefere WAV: %s", e)
|
||
audio_bytes = wav_bytes
|
||
content_type = "audio/wav"
|
||
else:
|
||
audio_bytes = wav_bytes
|
||
content_type = "audio/wav"
|
||
|
||
total_time = time.time() - start
|
||
log.info(
|
||
"Synthese: %d Zeichen, %s, %.1fs Audio, %.1fs Gesamt",
|
||
len(text),
|
||
output_format,
|
||
duration,
|
||
total_time,
|
||
)
|
||
return audio_bytes, content_type
|
||
|
||
def health(self) -> dict:
|
||
"""Liefert Health-Status."""
|
||
return {
|
||
"ready": self.status == "ready",
|
||
"status": self.status,
|
||
"model": self.model_name,
|
||
"voices": self.voices,
|
||
"error": self.error,
|
||
}
|
||
|
||
|
||
# Globale Worker-Instanz
|
||
worker = XTTSWorker()
|
||
|
||
|
||
class Handler(BaseHTTPRequestHandler):
|
||
server_version = "XTTSWorker/1.0"
|
||
timeout = 60
|
||
|
||
def do_GET(self):
|
||
if self.path == "/status":
|
||
self._send_json(200, worker.health())
|
||
else:
|
||
self._send_json(404, {"error": "not found"})
|
||
|
||
def do_POST(self):
|
||
if self.path == "/tts":
|
||
self._tts()
|
||
else:
|
||
self._send_json(404, {"error": "not found"})
|
||
|
||
def _tts(self):
|
||
try:
|
||
length = int(self.headers.get("Content-Length") or 0)
|
||
body = self.rfile.read(length)
|
||
data = json.loads(body)
|
||
except (ValueError, json.JSONDecodeError) as e:
|
||
self._send_json(400, {"error": f"Invalid JSON: {e}"})
|
||
return
|
||
|
||
text = data.get("text", "")
|
||
if not text:
|
||
self._send_json(400, {"error": "No text provided"})
|
||
return
|
||
|
||
voice = data.get("voice", DEFAULT_VOICE)
|
||
speed = float(data.get("speed", 1.0))
|
||
fmt = data.get("format", "mp3")
|
||
|
||
try:
|
||
audio_bytes, content_type = worker.synthesize(
|
||
text=text,
|
||
voice=voice,
|
||
speed=speed,
|
||
output_format=fmt,
|
||
)
|
||
self.send_response(200)
|
||
self.send_header("Content-Type", content_type)
|
||
self.send_header("Content-Length", str(len(audio_bytes)))
|
||
self.send_header("Connection", "close")
|
||
self.end_headers()
|
||
self.wfile.write(audio_bytes)
|
||
except Exception as e:
|
||
log.error("Synthese-Fehler: %s", e)
|
||
self._send_json(500, {"error": str(e)})
|
||
|
||
def _send_json(self, code: int, payload: dict) -> None:
|
||
body = json.dumps(payload).encode()
|
||
self.send_response(code)
|
||
self.send_header("Content-Type", "application/json")
|
||
self.send_header("Content-Length", str(len(body)))
|
||
self.send_header("Connection", "close")
|
||
self.end_headers()
|
||
self.wfile.write(body)
|
||
|
||
def log_message(self, format, *args):
|
||
# Logging nach stdout (journald)
|
||
log.info("%s - %s", self.address_string(), format % args)
|
||
|
||
|
||
def main():
|
||
# Modell laden
|
||
try:
|
||
worker.load()
|
||
except Exception as e:
|
||
log.error("Konnte Modell nicht laden: %s", e)
|
||
sys.exit(1)
|
||
|
||
log.info("Worker bereit auf Port %d", PORT)
|
||
|
||
server = ThreadingHTTPServer(("0.0.0.0", PORT), Handler)
|
||
server.daemon_threads = True
|
||
try:
|
||
server.serve_forever()
|
||
except KeyboardInterrupt:
|
||
pass
|
||
finally:
|
||
server.server_close()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|