diff options
| author | historia <historiavg@proton.me> | 2026-08-18 23:27:42 -0400 |
|---|---|---|
| committer | historia <historiavg@proton.me> | 2026-08-18 23:27:42 -0400 |
| commit | f3e21980320c1708ff17cc6f699a9aa4758accdf (patch) | |
| tree | 596ff71ba68600bb24a5a9d88424951c77e02808 /converter/tts.py | |
| parent | a97c0506f8b20cdc5ed8a11892ef10a9fc1938ef (diff) | |
| download | tts-audiobook-generator-f3e21980320c1708ff17cc6f699a9aa4758accdf.tar.gz | |
feat: support for faster-qwen3-tts backend server
Diffstat (limited to 'converter/tts.py')
| -rw-r--r-- | converter/tts.py | 314 |
1 files changed, 229 insertions, 85 deletions
diff --git a/converter/tts.py b/converter/tts.py index b6049a1..cb09936 100644 --- a/converter/tts.py +++ b/converter/tts.py @@ -1,16 +1,27 @@ -"""Client wrapper for the Qwen3-TTS Gradio demos (custom voice / voice clone).""" +"""Client wrappers for the TTS backends. + +QwenTTSClient talks to the Qwen3-TTS Gradio demos (custom voice / voice clone). +FasterTTSClient talks to the OpenAI-compatible server from the +faster-qwen3-tts repository (voice cloning only; the reference voice is +configured server-side — see the "Faster backend" section of the README). +""" import contextlib import io +import json import logging import shutil import sys import threading import time +import urllib.error +import urllib.request +import wave from pathlib import Path -from typing import Any, Dict, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple from . import config +from .chunking import split_into_chunks logger = logging.getLogger(__name__) @@ -47,7 +58,105 @@ def normalize_language(value: Optional[str]) -> str: ) -class QwenTTSClient: +def transcribe_reference_audio(audio_path: str, model_name: str = "base") -> Optional[str]: + """Transcribe reference audio locally using an optional Whisper backend. + + The current qwen-tts demo does not expose a transcription endpoint, so + transcription is done client-side when a Whisper package is available. + Returns None if no backend is installed. + """ + for backend in ("faster_whisper", "whisper"): + try: + if backend == "faster_whisper": + from faster_whisper import WhisperModel + model = WhisperModel(model_name, device="cpu", compute_type="int8") + segments, _ = model.transcribe(audio_path) + text = " ".join(seg.text.strip() for seg in segments).strip() + else: + import whisper + model = whisper.load_model(model_name) + result = model.transcribe(audio_path) + text = (result.get("text") or "").strip() + if text: + logger.info("Transcription complete via %s: %s", backend, text) + return text + except ImportError: + continue + except Exception as exc: + logger.warning("%s transcription failed: %s", backend, exc) + logger.warning("No Whisper backend available; transcription skipped.") + return None + + +class _BaseTTSClient: + """Shared chunk retry logic, heartbeat, and chunk file bookkeeping.""" + + def generate_chunk(self, text: str, chunk_num: int) -> Optional[str]: + """Generate one audio chunk; returns its path in the chunks folder.""" + raise NotImplementedError + + def _chunk_path(self, chunk_num: int, suffix: str) -> Path: + """Resolve the target path for a chunk, removing stale files first. + + Any stale chunk file for this index is removed so a retry or extension + change can never leave two files matching chunk_NNNN.*. + """ + for stale in config.CHUNKS_FOLDER.glob(f"chunk_{chunk_num:04d}.*"): + try: + stale.unlink() + except OSError as exc: + logger.debug("Could not remove stale chunk file %s: %s", stale, exc) + return config.CHUNKS_FOLDER / f"chunk_{chunk_num:04d}{suffix}" + + def process_chunk_with_retry(self, chunk_num: int, text: str) -> Optional[Path]: + """Process a chunk with retry logic and rate limiting. + + Returns the generated chunk file's path, or None when all attempts + failed. + """ + # Small delay between chunks to avoid rate limiting (only if not first chunk) + if chunk_num > 1: + time.sleep(config.MIN_DELAY_BETWEEN_CHUNKS) + + for attempt in range(config.MAX_RETRIES): + try: + result = self.generate_chunk(text, chunk_num) + if result and Path(result).exists(): + return Path(result) + logger.warning("Chunk %d attempt %d failed", chunk_num, attempt + 1) + except Exception as exc: + logger.warning("Chunk %d attempt %d error: %s", chunk_num, attempt + 1, exc) + + if attempt < config.MAX_RETRIES - 1: + sleep_time = 5 + (2 ** attempt) + logger.info("Waiting %ds before retry...", sleep_time) + time.sleep(sleep_time) + + logger.error("Chunk %d failed after %d attempts", chunk_num, config.MAX_RETRIES) + return None + + @contextlib.contextmanager + def _chunk_heartbeat(self, chunk_num: int): + """Print a periodic "still working" message while a chunk generates.""" + stop = threading.Event() + + def _beat(): + start = time.time() + while not stop.wait(config.HEARTBEAT_INTERVAL_SECONDS): + elapsed = time.time() - start + print(f"[...] Chunk {chunk_num} still generating — " + f"{int(elapsed // 60)}m {int(elapsed % 60)}s elapsed", flush=True) + + thread = threading.Thread(target=_beat, daemon=True) + thread.start() + try: + yield + finally: + stop.set() + thread.join() + + +class QwenTTSClient(_BaseTTSClient): """Generates audio chunks through a Qwen3-TTS Gradio server.""" def __init__(self, voice_mode: str = "custom_voice", voice_clone_ref_audio: Optional[str] = None, @@ -166,33 +275,8 @@ class QwenTTSClient: # ------------------------------------------------------------------ def transcribe_audio(self, audio_path: str) -> Optional[str]: - """Transcribe reference audio locally using an optional Whisper backend. - - The current qwen-tts demo does not expose a transcription endpoint, so - transcription is done client-side when a Whisper package is available. - Returns None if no backend is installed. - """ - for backend in ("faster_whisper", "whisper"): - try: - if backend == "faster_whisper": - from faster_whisper import WhisperModel - model = WhisperModel("base", device="cpu", compute_type="int8") - segments, _ = model.transcribe(audio_path) - text = " ".join(seg.text.strip() for seg in segments).strip() - else: - import whisper - model = whisper.load_model("base") - result = model.transcribe(audio_path) - text = (result.get("text") or "").strip() - if text: - logger.info("Transcription complete via %s: %s", backend, text) - return text - except ImportError: - continue - except Exception as exc: - logger.warning("%s transcription failed: %s", backend, exc) - logger.warning("No Whisper backend available; transcription skipped.") - return None + """Transcribe reference audio locally using an optional Whisper backend.""" + return transcribe_reference_audio(audio_path) # ------------------------------------------------------------------ # Chunk generation @@ -222,14 +306,7 @@ class QwenTTSClient: raise RuntimeError(f"Generated audio file not found: {audio_path}") suffix = source.suffix or ".wav" - # Remove any stale chunk file for this index first so a retry or - # extension change can never leave two files matching chunk_NNNN.* - for stale in config.CHUNKS_FOLDER.glob(f"chunk_{chunk_num:04d}.*"): - try: - stale.unlink() - except OSError as exc: - logger.debug("Could not remove stale chunk file %s: %s", stale, exc) - output_path = config.CHUNKS_FOLDER / f"chunk_{chunk_num:04d}{suffix}" + output_path = self._chunk_path(chunk_num, suffix) shutil.copy2(source, output_path) logger.debug("Chunk %d generated successfully", chunk_num) @@ -239,53 +316,6 @@ class QwenTTSClient: logger.error("Qwen chunk processing failed for chunk %d: %s", chunk_num, exc) return None - def process_chunk_with_retry(self, chunk_num: int, text: str) -> Optional[Path]: - """Process a chunk with retry logic and rate limiting. - - Returns the generated chunk file's path, or None when all attempts - failed. - """ - # Small delay between chunks to avoid rate limiting (only if not first chunk) - if chunk_num > 1: - time.sleep(config.MIN_DELAY_BETWEEN_CHUNKS) - - for attempt in range(config.MAX_RETRIES): - try: - result = self.generate_chunk(text, chunk_num) - if result and Path(result).exists(): - return Path(result) - logger.warning("Chunk %d attempt %d failed", chunk_num, attempt + 1) - except Exception as exc: - logger.warning("Chunk %d attempt %d error: %s", chunk_num, attempt + 1, exc) - - if attempt < config.MAX_RETRIES - 1: - sleep_time = 5 + (2 ** attempt) - logger.info("Waiting %ds before retry...", sleep_time) - time.sleep(sleep_time) - - logger.error("Chunk %d failed after %d attempts", chunk_num, config.MAX_RETRIES) - return None - - @contextlib.contextmanager - def _chunk_heartbeat(self, chunk_num: int): - """Print a periodic "still working" message while a chunk generates.""" - stop = threading.Event() - - def _beat(): - start = time.time() - while not stop.wait(config.HEARTBEAT_INTERVAL_SECONDS): - elapsed = time.time() - start - print(f"[...] Chunk {chunk_num} still generating — " - f"{int(elapsed // 60)}m {int(elapsed % 60)}s elapsed", flush=True) - - thread = threading.Thread(target=_beat, daemon=True) - thread.start() - try: - yield - finally: - stop.set() - thread.join() - # ------------------------------------------------------------------ # API payloads # ------------------------------------------------------------------ @@ -363,3 +393,117 @@ class QwenTTSClient: payload[name] = value return self.clone_client.predict(**payload, api_name=clone_api) + + +class FasterTTSClient(_BaseTTSClient): + """Generates audio chunks through a faster-qwen3-tts server. + + Talks to the OpenAI-compatible server shipped in the faster-qwen3-tts + repository (examples/openai_server.py). The reference voice (ref audio, + ref text) and language are configured on the server itself via + --ref-audio/--ref-text or a --voices JSON file; this client only sends + text. Unlike the Gradio demo, the server performs one generation per + request, so long chunks are sub-chunked client-side. + """ + + def __init__(self, voice: Optional[str] = None, api_url: Optional[str] = None): + self.voice = voice or config.FASTER_TTS_VOICE + self.api_url = (api_url or config.FASTER_TTS_API_URL).rstrip("/") + self._check_health() + + def _check_health(self) -> None: + """Verify the server is reachable and its model is loaded.""" + url = f"{self.api_url}/health" + try: + with urllib.request.urlopen(url, timeout=10) as response: + payload = json.loads(response.read().decode("utf-8")) + except Exception as exc: + raise RuntimeError( + f"Faster TTS server not reachable at {url}: {exc}. " + "Start the faster-qwen3-tts OpenAI-compatible server first " + "(see the 'Faster backend' section of the README)." + ) from exc + if not payload.get("model_loaded"): + raise RuntimeError( + "The faster TTS server is running but its model is not loaded yet; " + "wait for model download and startup to finish, then retry." + ) + print(f"[OK] Connected to faster TTS API at {self.api_url} (voice '{self.voice}')") + print(f"[INFO] The server silently falls back to its first configured voice if " + f"'{self.voice}' is not defined in its voice config (see README).") + + # ------------------------------------------------------------------ + # HTTP requests + # ------------------------------------------------------------------ + + def _request_pcm(self, text: str) -> bytes: + """POST one sub-chunk and return raw 16-bit mono PCM bytes.""" + url = f"{self.api_url}/v1/audio/speech" + payload = json.dumps({ + "model": "tts-1", + "input": text, + "voice": self.voice, + "response_format": "pcm", + }).encode("utf-8") + request = urllib.request.Request( + url, data=payload, headers={"Content-Type": "application/json"}, method="POST") + try: + with urllib.request.urlopen(request, timeout=config.FASTER_HTTP_TIMEOUT) as response: + pcm = response.read() + except urllib.error.HTTPError as exc: + detail = "" + try: + detail = exc.read().decode("utf-8", errors="replace")[:200] + except Exception: + pass + raise RuntimeError(f"Faster TTS server returned HTTP {exc.code}: {detail}") from exc + except urllib.error.URLError as exc: + raise RuntimeError(f"Faster TTS request failed: {exc.reason}") from exc + if not pcm: + raise RuntimeError("Faster TTS server returned empty audio") + return pcm + + def _request_pcm_with_retry(self, text: str, chunk_num: int, sub_num: int, + sub_total: int) -> bytes: + """Request one sub-chunk, retrying transient failures.""" + for attempt in range(config.FASTER_SUBCHUNK_RETRIES): + try: + return self._request_pcm(text) + except Exception as exc: + logger.warning("Chunk %d sub-chunk %d/%d attempt %d failed: %s", + chunk_num, sub_num, sub_total, attempt + 1, exc) + if attempt < config.FASTER_SUBCHUNK_RETRIES - 1: + time.sleep(2 + 2 * attempt) + raise RuntimeError(f"Sub-chunk {sub_num}/{sub_total} failed after " + f"{config.FASTER_SUBCHUNK_RETRIES} attempts") + + # ------------------------------------------------------------------ + # Chunk generation + # ------------------------------------------------------------------ + + def generate_chunk(self, text: str, chunk_num: int) -> Optional[str]: + """Generate one audio chunk; returns its path in the chunks folder.""" + try: + sub_chunks = split_into_chunks(text, max_words=config.FASTER_SUBCHUNK_WORDS) + if not sub_chunks: + raise RuntimeError("No text to synthesize") + + pcm_parts: List[bytes] = [] + with self._chunk_heartbeat(chunk_num): + for sub_num, sub_text in enumerate(sub_chunks, 1): + pcm_parts.append(self._request_pcm_with_retry( + sub_text, chunk_num, sub_num, len(sub_chunks))) + + output_path = self._chunk_path(chunk_num, ".wav") + with wave.open(str(output_path), "wb") as wav_file: + wav_file.setnchannels(1) + wav_file.setsampwidth(2) + wav_file.setframerate(config.FASTER_TTS_SAMPLE_RATE) + wav_file.writeframes(b"".join(pcm_parts)) + + logger.debug("Chunk %d generated (%d sub-chunks)", chunk_num, len(sub_chunks)) + return str(output_path) + + except Exception as exc: + logger.error("Faster chunk processing failed for chunk %d: %s", chunk_num, exc) + return None |
