refactor: extract TTSContext dataclass for normalization parameters (#5)

Bundles pronunciation_rules, heteronym_rules, normalization_overrides,
usage_counter, and split_pattern into a single TTSContext dataclass.
Both UIs create it once and use tts_context.normalize() instead of
threading 5 separate parameters through prepare_text_for_tts calls.

Tests: 1253 passed
This commit is contained in:
Artem Akymenko
2026-07-20 09:00:55 +00:00
parent 8144a7a507
commit c4870eece6
3 changed files with 52 additions and 23 deletions
+29
View File
@@ -5,10 +5,14 @@ and the comprehensive ``prepare_text_for_tts`` that chains all three normalizati
stages used during conversion: heteronym rules → pronunciation rules → pipeline stages used during conversion: heteronym rules → pronunciation rules → pipeline
normalization. The latter is the single entry point that both the Web UI and normalization. The latter is the single entry point that both the Web UI and
PyQt Desktop GUI should use. PyQt Desktop GUI should use.
Also provides ``TTSContext`` — a dataclass bundling all pre-compiled normalization
resources so they can be created once and passed as a single object.
""" """
from __future__ import annotations from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Dict, List, Mapping, Optional from typing import Any, Dict, List, Mapping, Optional
from abogen.kokoro_text_normalization import ( from abogen.kokoro_text_normalization import (
@@ -24,6 +28,31 @@ from abogen.normalization_settings import (
_BASE_APOSTROPHE_CONFIG = ApostropheConfig() _BASE_APOSTROPHE_CONFIG = ApostropheConfig()
@dataclass
class TTSContext:
"""Bundles pre-compiled normalization resources for TTS processing.
Created once per conversion job and passed to ``prepare_text_for_tts``
instead of threading 5 separate parameters.
"""
split_pattern: str = r"(?<=[.!?\-])\s+"
pronunciation_rules: Optional[List[Dict[str, Any]]] = None
heteronym_rules: Optional[List[Dict[str, Any]]] = None
normalization_overrides: Optional[Mapping[str, Any]] = None
usage_counter: Dict[str, int] = field(default_factory=dict)
def normalize(self, text: str) -> str:
"""Shorthand: normalize text using this context's compiled rules."""
return prepare_text_for_tts(
text,
heteronym_rules=self.heteronym_rules,
pronunciation_rules=self.pronunciation_rules,
normalization_overrides=self.normalization_overrides,
usage_counter=self.usage_counter,
)
def normalize_text_for_pipeline( def normalize_text_for_pipeline(
text: str, text: str,
*, *,
+9 -11
View File
@@ -47,7 +47,7 @@ from abogen.domain.audio_buffer import (
from abogen.domain.subtitle_generation import process_subtitle_tokens from abogen.domain.subtitle_generation import process_subtitle_tokens
from abogen.domain.voice_loader import VoiceCache, 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.progress import calc_etr_str
from abogen.domain.normalization import prepare_text_for_tts from abogen.domain.normalization import TTSContext
from abogen.domain.pronunciation import ( from abogen.domain.pronunciation import (
compile_pronunciation_rules, compile_pronunciation_rules,
compile_heteronym_sentence_rules, compile_heteronym_sentence_rules,
@@ -566,15 +566,19 @@ class ConversionThread(QThread):
) )
# --- Compile normalization rules (heteronym + pronunciation) --- # --- Compile normalization rules (heteronym + pronunciation) ---
from abogen.domain.normalization import TTSContext
pronunciation_overrides = merge_pronunciation_overrides( pronunciation_overrides = merge_pronunciation_overrides(
getattr(self, "pronunciation_overrides", None), getattr(self, "pronunciation_overrides", None),
getattr(self, "manual_overrides", None), getattr(self, "manual_overrides", None),
) )
self._pronunciation_rules = compile_pronunciation_rules(pronunciation_overrides) self._tts_context = TTSContext(
self._heteronym_rules = compile_heteronym_sentence_rules( split_pattern=self.split_pattern,
pronunciation_rules=compile_pronunciation_rules(pronunciation_overrides),
heteronym_rules=compile_heteronym_sentence_rules(
getattr(self, "heteronym_overrides", None) getattr(self, "heteronym_overrides", None)
),
normalization_overrides=getattr(self, "normalization_overrides", None),
) )
self._usage_counter = {}
# --- Chapter splitting logic --- # --- Chapter splitting logic ---
chapters = parse_chapters_from_text(text, clean=False) chapters = parse_chapters_from_text(text, clean=False)
@@ -980,13 +984,7 @@ class ConversionThread(QThread):
for text_segment in text_segments: for text_segment in text_segments:
# Normalize text before TTS # Normalize text before TTS
try: try:
text_segment = prepare_text_for_tts( text_segment = self._tts_context.normalize(text_segment)
text_segment,
heteronym_rules=getattr(self, "_heteronym_rules", None),
pronunciation_rules=getattr(self, "_pronunciation_rules", None),
normalization_overrides=getattr(self, "normalization_overrides", None),
usage_counter=getattr(self, "_usage_counter", None),
)
except Exception as exc: except Exception as exc:
self.log_updated.emit( self.log_updated.emit(
(f"Warning: Text normalization failed: {exc}", "orange") (f"Warning: Text normalization failed: {exc}", "orange")
+13 -11
View File
@@ -70,7 +70,7 @@ from abogen.domain.pronunciation import (
apply_pronunciation_rules as _apply_pronunciation_rules, apply_pronunciation_rules as _apply_pronunciation_rules,
merge_pronunciation_overrides as _merge_pronunciation_overrides, merge_pronunciation_overrides as _merge_pronunciation_overrides,
) )
from abogen.domain.normalization import prepare_text_for_tts from abogen.domain.normalization import TTSContext
from abogen.domain.voice_resolution import ( from abogen.domain.voice_resolution import (
spec_to_voice_ids as _spec_to_voice_ids, spec_to_voice_ids as _spec_to_voice_ids,
job_voice_fallback as _job_voice_fallback, job_voice_fallback as _job_voice_fallback,
@@ -242,6 +242,14 @@ def run_conversion_job(job: Job) -> None:
f"Applying {count} pronunciation override{'s' if count != 1 else ''} during conversion.", f"Applying {count} pronunciation override{'s' if count != 1 else ''} during conversion.",
level="debug", level="debug",
) )
tts_context = TTSContext(
split_pattern=job_split_pattern,
pronunciation_rules=pronunciation_rules,
heteronym_rules=heteronym_sentence_rules,
normalization_overrides=getattr(job, "normalization_overrides", None),
usage_counter=usage_counter,
)
for override_entry in pronunciation_overrides or []: for override_entry in pronunciation_overrides or []:
if not isinstance(override_entry, Mapping): if not isinstance(override_entry, Mapping):
continue continue
@@ -415,16 +423,10 @@ def run_conversion_job(job: Job) -> None:
) -> int: ) -> int:
nonlocal processed_chars, current_time nonlocal processed_chars, current_time
if split_pattern is None: if split_pattern is None:
split_pattern = job_split_pattern split_pattern = tts_context.split_pattern
source_text = str(text or "") source_text = str(text or "")
try: try:
normalized = prepare_text_for_tts( normalized = tts_context.normalize(source_text)
source_text,
heteronym_rules=heteronym_sentence_rules,
pronunciation_rules=pronunciation_rules,
normalization_overrides=getattr(job, "normalization_overrides", None),
usage_counter=usage_counter,
)
except LLMClientError as exc: except LLMClientError as exc:
job.add_log(f"LLM normalization failed: {exc}", level="error") job.add_log(f"LLM normalization failed: {exc}", level="error")
raise raise
@@ -890,8 +892,8 @@ def run_conversion_job(job: Job) -> None:
"generate_epub3": job.generate_epub3, "generate_epub3": job.generate_epub3,
} }
if usage_counter: if tts_context.usage_counter:
_record_override_usage(job, usage_counter, override_token_map) _record_override_usage(job, tts_context.usage_counter, override_token_map)
if metadata_dir: if metadata_dir:
metadata_dir.mkdir(parents=True, exist_ok=True) metadata_dir.mkdir(parents=True, exist_ok=True)