"""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__) def speaker_display_name() -> str: """Return the Gradio display name for the configured custom speaker.""" return config.SPEAKER_DISPLAY_NAMES.get( config.CUSTOM_VOICE_SPEAKER.lower(), config.CUSTOM_VOICE_SPEAKER) def normalize_language(value: Optional[str]) -> str: """Normalize a user-provided language name to a Qwen3-TTS display name. Accepts the display names in config.TTS_LANGUAGES case-insensitively as well as the short aliases in config.TTS_LANGUAGE_ALIASES (ISO 639-1 codes and common shorthands). Raises ValueError for anything else, since the Qwen3-TTS demo silently falls back to "Auto" for unrecognized languages. """ if value is None: raise ValueError("Language must not be None") candidate = value.strip() if not candidate: raise ValueError("Language must not be empty") for name in config.TTS_LANGUAGES: if candidate.lower() == name.lower(): return name alias = config.TTS_LANGUAGE_ALIASES.get(candidate.lower()) if alias: return alias raise ValueError( f"Unknown language: {value!r}. Expected one of " f"{', '.join(config.TTS_LANGUAGES)} (or an alias: " f"{', '.join(sorted(config.TTS_LANGUAGE_ALIASES))})." ) 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, language: Optional[str] = None): if voice_mode not in config.VOICE_MODES: raise ValueError( f"Unknown voice mode: {voice_mode!r} (expected one of {config.VOICE_MODES})" ) 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 if language is None: language = (config.VOICE_CLONE_LANGUAGE if voice_mode == config.VOICE_MODE_CLONE else config.CUSTOM_VOICE_LANGUAGE) # Validate before connecting so bad values fail fast without a server. self.language = normalize_language(language) 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: api_url = config.VOICE_CLONE_API_URL if self.voice_mode == config.VOICE_MODE_CLONE else config.QWEN_API_URL try: if self.voice_mode == config.VOICE_MODE_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: raise RuntimeError( f"Qwen API initialization failed at {api_url}: {exc}. " "Make sure the Qwen Gradio server is running and reachable, and that your " "installed Qwen3-TTS version matches this converter's API expectations " "(voice clone requires the Base-model demo: Qwen/Qwen3-TTS-12Hz-1.7B-Base)." ) from exc 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 --transcription "..." for higher-quality in-context cloning.') else: print(f"[OK] Reference text:\n{self.voice_clone_ref_text}") 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) 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 == config.VOICE_MODE_CUSTOM: with self._chunk_heartbeat(chunk_num): result = self._generate_custom_voice(text) elif self.voice_mode == config.VOICE_MODE_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 isinstance(result, (tuple, list)) or not result: raise RuntimeError("Qwen API returned an invalid result") audio_path = result[0] # First element is the audio file path if not isinstance(audio_path, (str, Path)) or not audio_path: raise RuntimeError("Qwen API did not return an audio file path") source = Path(audio_path) if not source.exists(): 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}" shutil.copy2(source, 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) -> 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 # ------------------------------------------------------------------ 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=self.language, spk_disp=speaker_display_name(), instruct=config.CUSTOM_VOICE_INSTRUCT, ) else: payload = dict( text=text, language=self.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=self.language, ) else: payload = dict( ref_audio=self._ref_audio_payload(), ref_text=self.voice_clone_ref_text, target_text=text, language=self.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)