aboutsummaryrefslogtreecommitdiff
path: root/converter/tts.py
diff options
context:
space:
mode:
authorhistoria <historiavg@proton.me>2026-08-17 18:32:14 -0400
committerhistoria <historiavg@proton.me>2026-08-17 18:32:14 -0400
commita30bd534151f757dad38763ae479fcd465d2ba0d (patch)
treed2185a9e12abb13b76944c6d39b1bd64f5f70e24 /converter/tts.py
parent2bb8f721092fe2f4d6ae6e425181094c1f0da57a (diff)
downloadtts-audiobook-generator-a30bd534151f757dad38763ae479fcd465d2ba0d.tar.gz
refactor(converter): extract logic into package
Diffstat (limited to 'converter/tts.py')
-rw-r--r--converter/tts.py311
1 files changed, 311 insertions, 0 deletions
diff --git a/converter/tts.py b/converter/tts.py
new file mode 100644
index 0000000..ac47ecb
--- /dev/null
+++ b/converter/tts.py
@@ -0,0 +1,311 @@
+"""Client wrapper for the Qwen3-TTS Gradio demos (custom voice / voice clone)."""
+
+import contextlib
+import io
+import logging
+import shutil
+import sys
+import threading
+import time
+from pathlib import Path
+from typing import Any, Dict, Optional, Tuple
+
+from . import config
+
+logger = logging.getLogger(__name__)
+
+
+class QwenTTSClient:
+ """Generates audio chunks through a Qwen3-TTS Gradio server."""
+
+ def __init__(self, voice_mode: str = "custom_voice", voice_clone_ref_audio: Optional[str] = None,
+ voice_clone_ref_text: Optional[str] = None, skip_transcription: bool = False):
+ self.voice_mode = voice_mode
+ self.voice_clone_ref_audio = voice_clone_ref_audio
+ self.voice_clone_ref_text = (voice_clone_ref_text or "").strip()
+ self.skip_transcription = skip_transcription
+ self.client = None
+ self.api_info: Dict[str, Any] = {}
+ self.clone_client = None
+ self.clone_api_info: Dict[str, Any] = {}
+ self._ref_audio_filedata: Optional[Dict[str, Any]] = None
+ self._connect()
+
+ # ------------------------------------------------------------------
+ # Connection
+ # ------------------------------------------------------------------
+
+ def _connect(self) -> None:
+ try:
+ if self.voice_mode == "voice_clone":
+ # Voice clone uses the Base-model demo, which is a separate server
+ # from the CustomVoice demo (that one only exposes /run_instruct).
+ self._init_client(config.VOICE_CLONE_API_URL, clone=True)
+ print(f"[OK] Connected to Voice Clone API at {config.VOICE_CLONE_API_URL}")
+ self._resolve_reference_text()
+ else:
+ self._init_client(config.QWEN_API_URL, clone=False)
+ print("[OK] Connected to Qwen API")
+ except Exception as exc:
+ api_url = config.VOICE_CLONE_API_URL if self.voice_mode == "voice_clone" else config.QWEN_API_URL
+ print("[ERROR] Qwen API initialization failed!")
+ print(f"API endpoint: {api_url}")
+ print("Make sure:")
+ print("1. Qwen Gradio server is running")
+ print("2. The server is accessible at the configured URL")
+ print("3. The endpoint URL is correct")
+ print("4. Your installed Qwen3-TTS version matches this converter's API expectations")
+ print(" (voice clone requires the Base-model demo: Qwen/Qwen3-TTS-12Hz-1.7B-Base)")
+ print(f"Error: {exc}")
+ sys.exit(1)
+
+ def _resolve_reference_text(self) -> None:
+ """Resolve the reference transcript: explicit text, then local
+ transcription, then x-vector-only mode."""
+ if not self.voice_clone_ref_text and self.voice_clone_ref_audio:
+ if self.skip_transcription:
+ print("[INFO] Skipping reference audio transcription (--no-transcription).")
+ else:
+ print("[INFO] Transcribing reference audio for voice cloning...")
+ self.voice_clone_ref_text = self.transcribe_audio(self.voice_clone_ref_audio) or ""
+ if not self.voice_clone_ref_text:
+ print("[WARNING] No reference text available; using x-vector-only clone mode (lower quality).")
+ print(' Pass --voice-sample-text "..." for higher-quality in-context cloning.')
+ else:
+ print(f"[OK] Reference text: {self.voice_clone_ref_text[:100]}...")
+
+ def _init_client(self, url: str, clone: bool = False) -> None:
+ """Initialize a Gradio client and store its API metadata."""
+ from gradio_client import Client
+
+ logger.info("Connecting to Qwen API at %s...", url)
+ old_stdout = sys.stdout
+ sys.stdout = io.TextIOWrapper(io.BytesIO(), encoding="utf-8", errors="replace")
+ try:
+ try:
+ client = Client(url, httpx_kwargs={"timeout": config.API_TIMEOUT})
+ except TypeError:
+ # Older gradio_client versions don't support httpx_kwargs.
+ client = Client(url)
+ finally:
+ sys.stdout = old_stdout
+ if clone:
+ self.clone_client = client
+ self.clone_api_info = self._load_api_info(client)
+ else:
+ self.client = client
+ self.api_info = self._load_api_info(client)
+ logger.info("Connected to Qwen API")
+
+ @staticmethod
+ def _load_api_info(client) -> Dict[str, Any]:
+ """Load available API metadata from the Gradio app."""
+ try:
+ return client.view_api(return_format="dict")
+ except Exception as exc:
+ logger.warning("Unable to read API metadata: %s", exc)
+ return {}
+
+ def _resolve_api_name(self, *candidates: str, api_info: Optional[Dict[str, Any]] = None) -> str:
+ """Return the first available api_name from candidate list."""
+ info = api_info if api_info is not None else self.api_info
+ named_endpoints = info.get("named_endpoints", {})
+ for candidate in candidates:
+ if candidate in named_endpoints:
+ return candidate
+ return candidates[0]
+
+ def _endpoint_accepts_param(self, api_name: str, param_name: str,
+ api_info: Optional[Dict[str, Any]] = None) -> bool:
+ """Check whether endpoint input schema includes the given parameter."""
+ info = api_info if api_info is not None else self.api_info
+ endpoint = info.get("named_endpoints", {}).get(api_name, {})
+ parameters = endpoint.get("parameters", [])
+ return any(parameter.get("parameter_name") == param_name for parameter in parameters)
+
+ # ------------------------------------------------------------------
+ # Reference audio transcription (voice clone)
+ # ------------------------------------------------------------------
+
+ 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[:100])
+ 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
+
+ # ------------------------------------------------------------------
+ # 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:
+ if self.voice_mode == "custom_voice":
+ with self._chunk_heartbeat(chunk_num):
+ result = self._generate_custom_voice(text)
+ elif self.voice_mode == "voice_clone":
+ with self._chunk_heartbeat(chunk_num):
+ result = self._generate_voice_clone(text)
+ else:
+ raise ValueError(f"Unknown voice mode: {self.voice_mode}")
+
+ if not result or len(result) < 2:
+ raise RuntimeError("Qwen API returned invalid result")
+
+ audio_path = result[0] # First element is the audio file path
+ if not audio_path or not Path(audio_path).exists():
+ raise RuntimeError(f"Generated audio file not found: {audio_path}")
+
+ output_path = config.CHUNKS_FOLDER / f"chunk_{chunk_num:04d}.wav"
+ shutil.copy2(audio_path, output_path)
+
+ logger.debug("Chunk %d generated successfully", chunk_num)
+ return str(output_path)
+
+ except Exception as exc:
+ 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) -> bool:
+ """Process a chunk with retry logic and rate limiting."""
+ # 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 True
+ 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 False
+
+ @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()
+
+ # ------------------------------------------------------------------
+ # API payloads
+ # ------------------------------------------------------------------
+
+ def _generate_custom_voice(self, text: str) -> Tuple:
+ """Generate audio using CustomVoice mode."""
+ custom_api = self._resolve_api_name("/run_instruct", "/run_custom_voice", "/generate_custom_voice")
+ if custom_api == "/run_instruct":
+ payload = dict(
+ text=text,
+ lang_disp=config.CUSTOM_VOICE_LANGUAGE,
+ spk_disp=config.SPEAKER_DISPLAY_NAMES.get(
+ config.CUSTOM_VOICE_SPEAKER.lower(), config.CUSTOM_VOICE_SPEAKER),
+ instruct=config.CUSTOM_VOICE_INSTRUCT,
+ )
+ else:
+ payload = dict(
+ text=text,
+ language=config.CUSTOM_VOICE_LANGUAGE,
+ speaker=config.CUSTOM_VOICE_SPEAKER,
+ instruct=config.CUSTOM_VOICE_INSTRUCT,
+ )
+ if self._endpoint_accepts_param(custom_api, "model_id_cv"):
+ payload["model_id_cv"] = config.CUSTOM_VOICE_MODEL_ID
+ elif self._endpoint_accepts_param(custom_api, "model_size"):
+ payload["model_size"] = config.CUSTOM_VOICE_MODEL_SIZE
+
+ if self._endpoint_accepts_param(custom_api, "seed"):
+ payload["seed"] = config.CUSTOM_VOICE_SEED
+
+ return self.client.predict(**payload, api_name=custom_api)
+
+ def _ref_audio_payload(self) -> Dict[str, Any]:
+ """Gradio file payload for the reference audio (built once, reused)."""
+ if self._ref_audio_filedata is None:
+ from gradio_client import handle_file
+ self._ref_audio_filedata = handle_file(self.voice_clone_ref_audio)
+ return self._ref_audio_filedata
+
+ def _generate_voice_clone(self, text: str) -> Tuple:
+ """Generate audio using Voice Clone mode."""
+ if not Path(self.voice_clone_ref_audio).exists():
+ raise FileNotFoundError(f"Reference audio not found: {self.voice_clone_ref_audio}")
+
+ if self.clone_client is None:
+ raise RuntimeError("Voice Clone client is not initialized. Is the Base-model demo running?")
+
+ clone_api = self._resolve_api_name("/run_voice_clone", "/generate_voice_clone",
+ api_info=self.clone_api_info)
+ use_xvector = config.VOICE_CLONE_USE_XVECTOR_ONLY or not self.voice_clone_ref_text
+
+ if clone_api == "/run_voice_clone":
+ payload = dict(
+ ref_aud=self._ref_audio_payload(),
+ ref_txt=self.voice_clone_ref_text,
+ use_xvec=use_xvector,
+ text=text,
+ lang_disp=config.VOICE_CLONE_LANGUAGE,
+ )
+ else:
+ payload = dict(
+ ref_audio=self._ref_audio_payload(),
+ ref_text=self.voice_clone_ref_text,
+ target_text=text,
+ language=config.VOICE_CLONE_LANGUAGE,
+ use_xvector_only=use_xvector,
+ )
+ optional_params = {
+ "model_size": config.VOICE_CLONE_MODEL_SIZE,
+ "max_chunk_chars": config.VOICE_CLONE_MAX_CHUNK_CHARS,
+ "chunk_gap": config.VOICE_CLONE_CHUNK_GAP,
+ "seed": config.VOICE_CLONE_SEED,
+ }
+ for name, value in optional_params.items():
+ if self._endpoint_accepts_param(clone_api, name, api_info=self.clone_api_info):
+ payload[name] = value
+
+ return self.clone_client.predict(**payload, api_name=clone_api)