diff --git a/abogen/domain/voice_loader.py b/abogen/domain/voice_loader.py index 3f7c9a8..5d86f99 100644 --- a/abogen/domain/voice_loader.py +++ b/abogen/domain/voice_loader.py @@ -62,10 +62,10 @@ def resolve_voice( # Check cache first if cache and cache.contains(voice_spec): return cache.get(voice_spec) - + # Load voice if "*" in voice_spec: - if pipeline is None: + if pipeline is None or not hasattr(pipeline, "load_single_voice"): return voice_spec loaded_voice = get_new_voice(pipeline, voice_spec, use_gpu) else: @@ -82,35 +82,43 @@ def load_voice_cached( voice_name: str, pipeline: Any, use_gpu: bool, - cache: Optional[Dict[str, Any]] = None, + cache: Any = None, ) -> Any: """Load voice with caching (compatibility wrapper for PyQt). - + This function maintains backward compatibility with the PyQt interface while using the unified voice loading logic. - + Args: voice_name: Voice name or formula string. pipeline: TTS pipeline instance. use_gpu: Whether to use GPU. - cache: Optional dict to use as cache (instead of VoiceCache). - + cache: Optional VoiceCache or dict to use as cache. + Returns: Loaded voice tensor or voice name string. """ - # Use dict cache if provided (for backward compatibility) + # Check cache (supports both VoiceCache and plain dict) if cache is not None: - if voice_name in cache: + if isinstance(cache, VoiceCache): + if cache.contains(voice_name): + return cache.get(voice_name) + elif voice_name in cache: return cache[voice_name] - + # Load voice if "*" in voice_name: + if pipeline is None or not hasattr(pipeline, "load_single_voice"): + return voice_name loaded_voice = get_new_voice(pipeline, voice_name, use_gpu) else: loaded_voice = voice_name - + # Cache it if cache is not None: - cache[voice_name] = loaded_voice - + if isinstance(cache, VoiceCache): + cache.set(voice_name, loaded_voice) + else: + cache[voice_name] = loaded_voice + return loaded_voice diff --git a/abogen/pyqt/conversion.py b/abogen/pyqt/conversion.py index 340d0f9..a9302de 100644 --- a/abogen/pyqt/conversion.py +++ b/abogen/pyqt/conversion.py @@ -44,7 +44,7 @@ from abogen.domain.audio_buffer import ( SAMPLE_RATE, ) from abogen.domain.subtitle_generation import process_subtitle_tokens -from abogen.domain.voice_loader import load_voice_cached, resolve_voice +from abogen.domain.voice_loader import VoiceCache, load_voice_cached, resolve_voice from abogen.domain.progress import calc_etr_str from abogen.domain.normalization import prepare_text_for_tts from abogen.domain.pronunciation import ( @@ -295,7 +295,7 @@ class ConversionThread(QThread): self.use_spacy_segmentation = True # Default, will be overridden from GUI # Set split pattern based on language and subtitle mode self.split_pattern = get_split_pattern(lang_code, subtitle_mode) - self.voice_cache = {} # Cache for loaded voices + self.voice_cache = VoiceCache() # Cache for loaded voices def load_voice_cached(self, voice_name, tts): """Load voice with caching to avoid reloading same voice. diff --git a/abogen/webui/conversion_runner.py b/abogen/webui/conversion_runner.py index a25f75c..18b21a7 100644 --- a/abogen/webui/conversion_runner.py +++ b/abogen/webui/conversion_runner.py @@ -119,6 +119,7 @@ from abogen.domain.audio_buffer import ( from abogen.domain.audio_sink import AudioSink, open_audio_sink from abogen.domain.pipeline_factory import PipelinePool from abogen.domain.conversion_engine import run_tts_segment_loop, process_and_write_subtitles, SegmentStats +from abogen.domain.voice_loader import VoiceCache, resolve_voice from abogen.domain.voice_utils import resolve_voice_target as _resolve_voice_target @@ -216,11 +217,11 @@ def run_conversion_job(job: Job) -> None: if provider == "kokoro": kokoro_backend = pipeline_pool.get("kokoro", job.language, job.use_gpu, job=job) - choice = _resolve_voice(kokoro_backend, resolved, job.use_gpu) + choice = resolve_voice(resolved, kokoro_backend, job.use_gpu, cache=voice_cache) else: choice = resolved - voice_cache[cache_key] = choice + voice_cache.set(cache_key, choice) return provider, resolved, choice, speed, steps extraction = extract_from_path(job.stored_path) @@ -360,7 +361,7 @@ def run_conversion_job(job: Job) -> None: chapter_dir.mkdir(parents=True, exist_ok=True) base_voice_spec = _job_voice_fallback(job) - voice_cache: Dict[str, Any] = {} + voice_cache = VoiceCache() base_provider, base_voice_resolved, _, _ = _resolve_voice_target( base_voice_spec, normalized_profiles, job_voice=getattr(job, "voice", "M1"), @@ -368,7 +369,7 @@ def run_conversion_job(job: Job) -> None: ) if base_provider == "kokoro" and base_voice_resolved and "*" not in base_voice_resolved: kokoro_backend = pipeline_pool.get("kokoro", job.language, job.use_gpu, job=job) - voice_cache[f"kokoro:{base_voice_resolved}"] = _resolve_voice(kokoro_backend, base_voice_resolved, job.use_gpu) + voice_cache.set(f"kokoro:{base_voice_resolved}", resolve_voice(base_voice_resolved, kokoro_backend, job.use_gpu)) processed_chars = 0 current_time = 0.0 etr_start_time = time.time() @@ -550,8 +551,7 @@ def run_conversion_job(job: Job) -> None: voice_choice = voice_cache.get(chapter_cache_key) if voice_choice is None: kokoro_backend = pipeline_pool.get("kokoro", job.language, job.use_gpu, job=job) - voice_choice = _resolve_voice(kokoro_backend, chapter_voice_resolved, job.use_gpu) - voice_cache[chapter_cache_key] = voice_choice + voice_choice = resolve_voice(chapter_voice_resolved, kokoro_backend, job.use_gpu, cache=voice_cache) else: voice_choice = chapter_voice_resolved @@ -694,12 +694,12 @@ def run_conversion_job(job: Job) -> None: chunk_voice_choice = voice_cache.get(chunk_cache_key) if chunk_voice_choice is None: kokoro_backend = pipeline_pool.get("kokoro", job.language, job.use_gpu, job=job) - chunk_voice_choice = _resolve_voice( - kokoro_backend, + chunk_voice_choice = resolve_voice( chunk_voice_resolved, + kokoro_backend, job.use_gpu, + cache=voice_cache, ) - voice_cache[chunk_cache_key] = chunk_voice_choice else: chunk_voice_choice = chunk_voice_resolved @@ -1066,14 +1066,6 @@ def _prepare_project_layout(job: Job, base_dir: Path) -> tuple[Path, Path, Path, -def _resolve_voice(pipeline, voice_spec: str, use_gpu: bool): - if "*" in voice_spec: - if pipeline is None or not hasattr(pipeline, "load_single_voice"): - return voice_spec - return get_new_voice(pipeline, voice_spec, use_gpu) - return voice_spec - - def _make_canceller(job: Job) -> Callable[[], None]: def _cancel() -> None: diff --git a/abogen/webui/debug_tts_runner.py b/abogen/webui/debug_tts_runner.py index 374a8f9..b2b735f 100644 --- a/abogen/webui/debug_tts_runner.py +++ b/abogen/webui/debug_tts_runner.py @@ -14,7 +14,8 @@ from abogen.kokoro_text_normalization import normalize_for_pipeline from abogen.normalization_settings import build_apostrophe_config from abogen.text_extractor import extract_from_path from abogen.voice_cache import ensure_voice_assets -from abogen.webui.conversion_runner import SAMPLE_RATE, _select_device, _to_float32, _resolve_voice, _spec_to_voice_ids +from abogen.webui.conversion_runner import SAMPLE_RATE, _select_device, _to_float32, _spec_to_voice_ids +from abogen.domain.voice_loader import resolve_voice from abogen.domain.split_pattern import get_split_pattern from abogen.tts_plugin.utils import create_pipeline @@ -176,7 +177,7 @@ def run_debug_tts_wavs( pass pipeline = _load_pipeline(language, use_gpu) - voice_choice = _resolve_voice(pipeline, voice_spec, use_gpu) + voice_choice = resolve_voice(voice_spec, pipeline, use_gpu) apostrophe_config = build_apostrophe_config(settings=settings) normalization_settings = dict(settings)