""" OpenAI-compatible TTS server for Qwen3-TTS CustomVoice models. Voices are defined in customvoice_voices.json as: { "voice_id": { "speaker": "Ryan", "language": "English", "instruct": "" } } The request body may also include "language" and "instruct" fields to override the configured defaults for a single generation. """ import argparse import asyncio import json import logging import queue import sys import threading from typing import Optional import numpy as np import uvicorn from fastapi import FastAPI, HTTPException from fastapi.responses import JSONResponse, Response, StreamingResponse from pydantic import BaseModel from contextlib import asynccontextmanager sys.path.append("/app") from faster_qwen3_tts.model import FasterQwen3TTS logger = logging.getLogger(__name__) logging.basicConfig(level=logging.INFO) app = FastAPI() tts_model: FasterQwen3TTS = None voices: dict = {} default_voice: str = None SAMPLE_RATE = 24000 DEFAULT_MAX_NEW_TOKENS = 2048 _model_lock = threading.Lock() _load_model_kwargs = None aligner_model = None def _get_aligner(): global aligner_model if aligner_model is None: try: from qwen_asr import Qwen3ForcedAligner import torch except ImportError: raise HTTPException(status_code=500, detail="qwen-asr is not installed. Run: pip install qwen-asr") logger.info("Loading Qwen3-ForcedAligner-0.6B...") aligner_model = Qwen3ForcedAligner.from_pretrained( "Qwen/Qwen3-ForcedAligner-0.6B", dtype=torch.bfloat16, device_map="cuda" ) logger.info("Aligner loaded.") return aligner_model @asynccontextmanager async def lifespan(app: FastAPI): loop = asyncio.get_event_loop() loop.run_in_executor(None, _do_load_and_warmup) yield def _do_load_and_warmup(): global tts_model, SAMPLE_RATE import torch args = _load_model_kwargs try: logger.info("Loading CustomVoice model %s ...", args.model) model = FasterQwen3TTS.from_pretrained( args.model, device=args.device, dtype=torch.bfloat16, attn_implementation="sdpa", max_seq_len=args.max_seq_len, ) SAMPLE_RATE = model.sample_rate logger.info("Model ready. Sample rate: %d Hz", SAMPLE_RATE) # Warmup logger.info("Warming up CUDA graphs (first request will be fast)...") try: for _ in model.generate_custom_voice_streaming( text="Warmup.", speaker="Ryan", language="English" ): pass logger.info("CUDA warmup complete — server ready.") except Exception as exc: logger.warning("Warmup failed (non-fatal): %s", exc) tts_model = model except Exception as exc: logger.error("Failed to load model: %s", exc) app = FastAPI(lifespan=lifespan) class SpeechRequest(BaseModel): model: str = "tts-1" input: str voice: str = "Ryan" response_format: str = "wav" # wav | pcm | mp3 | zip speed: float = 1.0 language: Optional[str] = None instruct: Optional[str] = None max_new_tokens: Optional[int] = None def _to_pcm16(audio: np.ndarray) -> bytes: return (audio * 32767).clip(-32768, 32767).astype(np.int16).tobytes() def _wav_header(sample_rate: int) -> bytes: import struct return struct.pack( "<4sI4s4sIHHIIHH4sI", b"RIFF", 0xFFFFFFFF, b"WAVE", b"fmt ", 16, 1, 1, sample_rate, sample_rate * 2, 2, 16, b"data", 0xFFFFFFFF, ) def _to_mp3_bytes(audio: np.ndarray, sr: int) -> bytes: import io from pydub import AudioSegment pcm = _to_pcm16(audio) seg = AudioSegment(pcm, frame_rate=sr, sample_width=2, channels=1) buf = io.BytesIO() seg.export(buf, format="mp3") return buf.getvalue() def resolve_voice(name: str) -> dict: cfg = voices.get(name) if cfg: return cfg if default_voice and default_voice in voices: logger.warning("Voice %r not found, falling back to %r", name, default_voice) return voices[default_voice] raise HTTPException(status_code=404, detail=f"Voice {name!r} not found") def _request_generation_params(req: SpeechRequest, voice_cfg: dict) -> dict: return { "text": req.input, "speaker": voice_cfg.get("speaker") or req.voice, "language": req.language or voice_cfg.get("language", "Auto"), "instruct": req.instruct if req.instruct is not None else voice_cfg.get("instruct") or None, "max_new_tokens": req.max_new_tokens or int(voice_cfg.get("max_new_tokens", DEFAULT_MAX_NEW_TOKENS)), } async def _stream_chunks(params: dict, speed: float): q: queue.Queue = queue.Queue() done = object() def producer(): process = None if speed != 1.0: import subprocess cmd = [ "ffmpeg", "-y", "-loglevel", "error", "-f", "s16le", "-ar", str(SAMPLE_RATE), "-ac", "1", "-i", "pipe:0", "-filter:a", f"atempo={speed}", "-f", "s16le", "-ar", str(SAMPLE_RATE), "-ac", "1", "pipe:1" ] process = subprocess.Popen(cmd, stdin=subprocess.PIPE, stdout=subprocess.PIPE) def ffmpeg_reader(): try: while True: out = process.stdout.read(4096) if not out: break q.put(out) except Exception as e: q.put(e) finally: q.put(done) import threading threading.Thread(target=ffmpeg_reader, daemon=True).start() try: with _model_lock: for chunk, _sr, _timing in tts_model.generate_custom_voice_streaming(**params): raw = _to_pcm16(chunk) if process: process.stdin.write(raw) process.stdin.flush() else: q.put(raw) except Exception as exc: q.put(exc) finally: if process: try: process.stdin.close() except Exception: pass else: q.put(done) import threading threading.Thread(target=producer, daemon=True).start() loop = asyncio.get_event_loop() while True: item = await loop.run_in_executor(None, q.get) if item is done: break if isinstance(item, Exception): raise item yield item @app.get("/health") async def health(): return {"status": "ok", "model_loaded": tts_model is not None} @app.post("/v1/audio/speech") async def create_speech(req: SpeechRequest): if tts_model is None: raise HTTPException(status_code=503, detail="Model not loaded") req.input = req.input.strip() if not req.input: raise HTTPException(status_code=400, detail="'input' text is empty") voice_cfg = resolve_voice(req.voice) params = _request_generation_params(req, voice_cfg) fmt = req.response_format.lower() content_types = {"wav": "audio/wav", "pcm": "audio/pcm", "mp3": "audio/mpeg", "zip": "application/zip"} if fmt not in content_types: raise HTTPException(status_code=400, detail=f"Unsupported format: {fmt!r}") if fmt in ("mp3", "zip"): loop = asyncio.get_event_loop() def generate(): with _model_lock: return tts_model.generate_custom_voice(**params) audio_arrays, sr = await loop.run_in_executor(None, generate) audio = audio_arrays[0] if audio_arrays else np.zeros(1, dtype=np.float32) if req.speed != 1.0: import subprocess cmd = [ "ffmpeg", "-y", "-loglevel", "error", "-f", "f32le", "-ar", str(sr), "-ac", "1", "-i", "pipe:0", "-filter:a", f"atempo={req.speed}", "-f", "f32le", "-ar", str(sr), "-ac", "1", "pipe:1" ] process = subprocess.Popen(cmd, stdin=subprocess.PIPE, stdout=subprocess.PIPE) process.stdin.write(audio.tobytes()) process.stdin.close() out = process.stdout.read() audio = np.frombuffer(out, dtype=np.float32) if fmt == "zip": def _align(): aligner = _get_aligner() res = aligner.align(audio=(audio, sr), text=req.input, language=voice_cfg.get("language", "Auto")) import dataclasses return [dataclasses.asdict(x) for x in res] align_data = await loop.run_in_executor(None, _align) import zipfile import io mp3_bytes = _to_mp3_bytes(audio, sr) zip_buf = io.BytesIO() with zipfile.ZipFile(zip_buf, "w", zipfile.ZIP_DEFLATED) as zf: zf.writestr("audio.mp3", mp3_bytes) zf.writestr("timer.json", json.dumps(align_data, ensure_ascii=False)) return Response(content=zip_buf.getvalue(), media_type=content_types[fmt]) return Response(content=_to_mp3_bytes(audio, sr), media_type="audio/mpeg") async def audio_stream(): if fmt == "wav": yield _wav_header(SAMPLE_RATE) async for raw in _stream_chunks(params, req.speed): yield raw return StreamingResponse(audio_stream(), media_type=content_types[fmt]) _voice_list = None _models_response = None def _build_voice_list(): global _voice_list, _models_response _voice_list = [{"id": v, "object": "model", "created": 1686935002, "owned_by": "qwen"} for v in voices] _models_response = {"object": "list", "data": _voice_list} @app.get("/v1/models") async def list_models(): return _models_response @app.get("/v1/audio/voices") async def list_audio_voices(): return _models_response @app.get("/v1/audio/models") async def list_audio_models(): return _models_response @app.get("/speakers") async def get_speakers(): return list(voices.keys()) @app.options("/{path:path}") async def options_handler(path: str): return JSONResponse(content={"status": "ok"}) def main(): global voices, default_voice, SAMPLE_RATE, DEFAULT_MAX_NEW_TOKENS, _load_model_kwargs parser = argparse.ArgumentParser() parser.add_argument("--model", default="/models/Qwen3-TTS-CustomVoice") parser.add_argument("--voices", default="/config/customvoice_voices.json") parser.add_argument("--port", type=int, default=8000) parser.add_argument("--host", default="0.0.0.0") parser.add_argument("--device", default="cuda") parser.add_argument("--max-seq-len", type=int, default=2048) args = parser.parse_args() # Force a smaller max_seq_len to save VRAM and prevent OOM args.max_seq_len = 1024 DEFAULT_MAX_NEW_TOKENS = args.max_seq_len _load_model_kwargs = args with open(args.voices) as f: voices = json.load(f) default_voice = next(iter(voices), None) _build_voice_list() uvicorn.run(app, host=args.host, port=args.port, log_level="info") if __name__ == "__main__": main()