diff --git a/examples/openai_server.py b/examples/openai_server.py index 2199e14..d10065f 100644 --- a/examples/openai_server.py +++ b/examples/openai_server.py @@ -36,6 +36,7 @@ API usage: """ import argparse import asyncio +import hashlib import io import json import logging @@ -66,9 +67,46 @@ app = FastAPI(title="faster-qwen3-tts OpenAI-compatible API") tts_model = None voices: dict = {} +voices_file_path: Optional[str] = None +last_voices_mtime: float = 0.0 default_voice: Optional[str] = None SAMPLE_RATE = 24000 # updated once the model loads _model_lock = threading.Lock() # prevent concurrent GPU inference +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 + + +def _voice_seed(voice_name: str) -> int: + """Return a stable per-voice seed derived from the voice name. + + Used as the default when no explicit 'seed' is set in voices.json. + MD5 is used only for its stable byte output — not for security. + """ + return int(hashlib.md5(voice_name.encode()).hexdigest(), 16) % (2 ** 31) + + +def _seed_rng(seed: int) -> None: + """Seed PyTorch CPU and CUDA RNGs for reproducible sampling.""" + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + # --------------------------------------------------------------------------- # Request / response models @@ -79,8 +117,8 @@ class SpeechRequest(BaseModel): model: str = "tts-1" input: str voice: str = "alloy" - response_format: str = "wav" # wav | pcm | mp3 - speed: float = 1.0 # accepted but not yet applied + response_format: str = "wav" # wav | pcm | mp3 | zip + speed: float = 1.0 # scales audio tempo # --------------------------------------------------------------------------- @@ -145,6 +183,21 @@ def _to_mp3_bytes(pcm: np.ndarray, sample_rate: int) -> bytes: def resolve_voice(voice_name: str) -> dict: """Return voice config dict or fall back to default, else raise 400.""" + global voices, last_voices_mtime + voice_name = voice_name.strip() + + # Hot-reload voices.json if it was modified + if voices_file_path and os.path.exists(voices_file_path): + try: + current_mtime = os.path.getmtime(voices_file_path) + if current_mtime > last_voices_mtime: + with open(voices_file_path, "r", encoding="utf-8") as f: + voices = json.load(f) + last_voices_mtime = current_mtime + logger.info("Hot-reloaded %d voices from %s", len(voices), voices_file_path) + except Exception as e: + logger.warning("Failed to hot-reload voices.json: %s", e) + if voice_name in voices: return voices[voice_name] if default_voice and default_voice in voices: @@ -168,7 +221,47 @@ def resolve_voice(voice_name: str) -> dict: # --------------------------------------------------------------------------- -async def _stream_chunks(voice_cfg: dict, text: str) -> AsyncGenerator[bytes, None]: +def _load_voice_clone_prompt(voice_cfg: dict, voice_name: str, tts_model): + spk_emb_path = voice_cfg.get("speaker_embeddings") or voice_cfg.get("speaker embeddings") + if spk_emb_path and os.path.isfile(spk_emb_path): + try: + return torch.load(spk_emb_path, map_location="cpu", weights_only=False) + except Exception as e: + logger.error("Failed to load speaker embeddings from %s: %s", spk_emb_path, e) + + # Auto-generate if missing + ref_audio = voice_cfg.get("ref_audio") + if not ref_audio or not os.path.isfile(ref_audio): + return None + + logger.info("Precomputing and saving speaker embedding for voice %r...", voice_name) + try: + ref_text = voice_cfg.get("ref_text", "") + # generate prompt using the model's built-in helper + prompt_items = tts_model.model.create_voice_clone_prompt(ref_audio, [ref_text]) + vcp = tts_model.model._prompt_items_to_voice_clone_prompt(prompt_items) + + # save it to the speakers directory + # If the file path is already in voice_cfg but doesn't exist, use that, otherwise generate a path + if spk_emb_path and not os.path.exists(spk_emb_path) and spk_emb_path.endswith('.pt'): + pt_path = spk_emb_path + else: + pt_path = f"/config/speakers/{voice_name}.pt" + + os.makedirs(os.path.dirname(pt_path), exist_ok=True) + torch.save(vcp, pt_path) + logger.info("Saved speaker embedding to %s", pt_path) + + # update in memory so future requests skip generating + voice_cfg["speaker_embeddings"] = pt_path + + return vcp + except Exception as e: + logger.error("Failed to precompute speaker embedding: %s", e) + return None + + +async def _stream_chunks(voice_cfg: dict, text: str, voice_name: str, speed: float) -> AsyncGenerator[bytes, None]: """ Run generate_voice_clone_streaming in a background thread and yield raw PCM bytes for each chunk as they arrive. @@ -177,21 +270,63 @@ async def _stream_chunks(voice_cfg: dict, text: str) -> AsyncGenerator[bytes, No _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) + + threading.Thread(target=ffmpeg_reader, daemon=True).start() + try: with _model_lock: + _seed_rng(voice_cfg.get("seed", _voice_seed(voice_name))) for chunk, _sr, _timing in tts_model.generate_voice_clone_streaming( text=text, language=voice_cfg.get("language", "Auto"), - ref_audio=voice_cfg["ref_audio"], + ref_audio=voice_cfg.get("ref_audio"), ref_text=voice_cfg.get("ref_text", ""), chunk_size=voice_cfg.get("chunk_size", 12), - non_streaming_mode=False, + instruct=voice_cfg.get("instruct"), + voice_clone_prompt=_load_voice_clone_prompt(voice_cfg, voice_name, tts_model), + non_streaming_mode=True, + temperature=voice_cfg.get("temperature", 0.8), + top_k=voice_cfg.get("top_k", 50), + top_p=voice_cfg.get("top_p", 0.9), ): - q.put(chunk) + 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: - q.put(_DONE) + if process: + try: + process.stdin.close() + except Exception: + pass + else: + q.put(_DONE) thread = threading.Thread(target=producer, daemon=True) thread.start() @@ -203,7 +338,7 @@ async def _stream_chunks(voice_cfg: dict, text: str) -> AsyncGenerator[bytes, No break if isinstance(item, Exception): raise item - yield _to_pcm16(item) + yield item # --------------------------------------------------------------------------- @@ -220,7 +355,8 @@ async def health(): async def create_speech(req: SpeechRequest): if tts_model is None: raise HTTPException(status_code=503, detail="Model not loaded") - if not req.input.strip(): + 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) @@ -230,36 +366,77 @@ async def create_speech(req: SpeechRequest): "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"response_format {fmt!r} not supported. Use: wav, pcm, mp3", + detail=f"response_format {fmt!r} not supported. Use: wav, pcm, mp3, zip", ) content_type = _CONTENT_TYPES[fmt] - # --- MP3: generate all audio, then encode (non-streaming) --- - if fmt == "mp3": + # --- MP3 / ZIP: generate all audio, then encode (non-streaming) --- + if fmt in ("mp3", "zip"): loop = asyncio.get_event_loop() def _generate(): with _model_lock: + _seed_rng(voice_cfg.get("seed", _voice_seed(req.voice))) return tts_model.generate_voice_clone( text=req.input, language=voice_cfg.get("language", "Auto"), - ref_audio=voice_cfg["ref_audio"], + ref_audio=voice_cfg.get("ref_audio"), ref_text=voice_cfg.get("ref_text", ""), + instruct=voice_cfg.get("instruct"), + voice_clone_prompt=_load_voice_clone_prompt(voice_cfg, req.voice, tts_model), + non_streaming_mode=True, + temperature=voice_cfg.get("temperature", 0.8), + top_k=voice_cfg.get("top_k", 50), + top_p=voice_cfg.get("top_p", 0.9), ) 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 + 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_type) + return Response(content=_to_mp3_bytes(audio, sr), media_type=content_type) # --- WAV / PCM: stream chunks as they are generated --- async def audio_stream(): if fmt == "wav": yield _wav_header(SAMPLE_RATE) # stream with unknown data length - async for raw_chunk in _stream_chunks(voice_cfg, req.input): + async for raw_chunk in _stream_chunks(voice_cfg, req.input, req.voice, req.speed): yield raw_chunk return StreamingResponse(audio_stream(), media_type=content_type) @@ -306,16 +483,20 @@ def _parse_args(): p.add_argument("--host", default="0.0.0.0", help="Bind host (default: 0.0.0.0)") p.add_argument("--port", type=int, default=8000, help="Bind port (default: 8000)") p.add_argument("--device", default="cuda", help="Torch device (default: cuda)") + p.add_argument("--max-seq-len", type=int, default=4096, help="Max sequence length for CUDA graph static cache (default: 4096)") return p.parse_args() def main(): - global tts_model, voices, default_voice, SAMPLE_RATE + global tts_model, voices, voices_file_path, last_voices_mtime, default_voice, SAMPLE_RATE args = _parse_args() # Build voice registry if args.voices: + voices_file_path = args.voices + if os.path.exists(args.voices): + last_voices_mtime = os.path.getmtime(args.voices) with open(args.voices) as f: voices = json.load(f) default_voice = next(iter(voices)) @@ -344,6 +525,7 @@ def main(): args.model, device=args.device, dtype=torch.bfloat16, + max_seq_len=args.max_seq_len, ) SAMPLE_RATE = tts_model.sample_rate logger.info("Model ready. Sample rate: %d Hz", SAMPLE_RATE)