feat: Supertonic language + total_steps propagation

Language enum expanded from 9 to 33 languages:
- Added 24 new ISO 639-1 languages: AR, BG, CS, DA, DE, EL, ET, FI,
  HR, HU, ID, KO, LT, LV, NL, PL, RO, RU, SK, SL, SV, TR, UK, VI
- Updated display_name, is_cjk (added KO)

Supertonic language mapping (32 languages, no ZH):
- engine.py: _SUPERTONIC_LANG_MAP, engine_language(), supported_languages()
- __init__.py: create_engine() passes config.language to pipeline
- pipeline.py: __init__() accepts language, resolves to ISO code;
  __call__() passes lang= to TTS.synthesize()

total_steps propagation:
- tts_segments(): +total_steps param, conditionally passed to backend
- synthesize_text(): +total_steps param
- run_tts_segment_loop(): +total_steps param
- executor: all 5 synthesize_text() calls pass total_steps

Integration:
- pipeline_factory: create_pipeline_for_job() passes language to supertonic
- preview path: create_pipeline('supertonic', language=language)

Tests updated to accept total_steps in FakeBackend.__call__
This commit is contained in:
Artem Akymenko
2026-07-28 14:09:42 +03:00
parent 953bef1e71
commit 2b70b9ca45
14 changed files with 146 additions and 15 deletions
@@ -277,6 +277,7 @@ def execute_conversion(
backend=intro_backend, backend=intro_backend,
voice=intro_voice, voice=intro_voice,
speed=intro_speed or request.speed, speed=intro_speed or request.speed,
total_steps=intro_steps,
chapter_sink=None, chapter_sink=None,
preview_callback=lambda text: events.log(f" {text[:80]}"), preview_callback=lambda text: events.log(f" {text[:80]}"),
) )
@@ -347,6 +348,7 @@ def execute_conversion(
backend=intro_backend, backend=intro_backend,
voice=intro_voice, voice=intro_voice,
speed=intro_speed or request.speed, speed=intro_speed or request.speed,
total_steps=intro_steps,
chapter_sink=chapter_sink, chapter_sink=chapter_sink,
preview_callback=lambda text: events.log(f" Intro: {text[:80]}"), preview_callback=lambda text: events.log(f" Intro: {text[:80]}"),
) )
@@ -421,6 +423,7 @@ def execute_conversion(
seg_provider = chapter_provider seg_provider = chapter_provider
seg_voice = chapter_voice seg_voice = chapter_voice
seg_speed = chapter_speed seg_speed = chapter_speed
seg_steps = chapter_steps
seg_backend = chapter_backend seg_backend = chapter_backend
# Track voice for chapter marker # Track voice for chapter marker
@@ -452,6 +455,7 @@ def execute_conversion(
backend=seg_backend, backend=seg_backend,
voice=seg_voice, voice=seg_voice,
speed=seg_speed or request.speed, speed=seg_speed or request.speed,
total_steps=seg_steps,
chapter_sink=chapter_sink, chapter_sink=chapter_sink,
preview_callback=lambda text: events.log(f" {text[:80]}"), preview_callback=lambda text: events.log(f" {text[:80]}"),
split_pattern_override=active_split, split_pattern_override=active_split,
@@ -539,6 +543,7 @@ def execute_conversion(
backend=outro_backend, backend=outro_backend,
voice=outro_voice, voice=outro_voice,
speed=outro_speed or request.speed, speed=outro_speed or request.speed,
total_steps=outro_steps,
chapter_sink=None, chapter_sink=None,
preview_callback=lambda text: events.log(f" {text[:80]}"), preview_callback=lambda text: events.log(f" {text[:80]}"),
) )
+5
View File
@@ -61,6 +61,7 @@ def run_tts_segment_loop(
voice: Any, voice: Any,
speed: float, speed: float,
split_pattern: str, split_pattern: str,
total_steps: Optional[int] = None,
chapter_sink: Optional[AudioSink] = None, chapter_sink: Optional[AudioSink] = None,
preview_callback: Optional[Callable[[str], None]] = None, preview_callback: Optional[Callable[[str], None]] = None,
on_segment: Optional[Callable[[SegmentInfo], None]] = None, on_segment: Optional[Callable[[SegmentInfo], None]] = None,
@@ -74,6 +75,7 @@ def run_tts_segment_loop(
voice: Voice name/id for the backend. voice: Voice name/id for the backend.
speed: Speech speed multiplier. speed: Speech speed multiplier.
split_pattern: Regex pattern used by the TTS engine for sentence splitting. split_pattern: Regex pattern used by the TTS engine for sentence splitting.
total_steps: Inference quality steps (Supertonic only, ignored by Kokoro).
preview_callback: Called with a short preview string per segment. preview_callback: Called with a short preview string per segment.
on_segment: Called with a SegmentInfo for each segment *before* on_segment: Called with a SegmentInfo for each segment *before*
audio is written. Useful for callers that need per-segment audio is written. Useful for callers that need per-segment
@@ -95,6 +97,7 @@ def run_tts_segment_loop(
speed=speed, speed=speed,
split_pattern=split_pattern, split_pattern=split_pattern,
current_time=params.stats.current_time, current_time=params.stats.current_time,
total_steps=total_steps,
): ):
if params.check_cancel(): if params.check_cancel():
break break
@@ -212,6 +215,7 @@ def synthesize_text(
backend: Any, backend: Any,
voice: Any, voice: Any,
speed: float, speed: float,
total_steps: Optional[int] = None,
chapter_sink: Optional[AudioSink] = None, chapter_sink: Optional[AudioSink] = None,
preview_callback: Optional[Callable[[str], None]] = None, preview_callback: Optional[Callable[[str], None]] = None,
on_segment: Optional[Callable[[SegmentInfo], None]] = None, on_segment: Optional[Callable[[SegmentInfo], None]] = None,
@@ -229,6 +233,7 @@ def synthesize_text(
backend=backend, backend=backend,
voice=voice, voice=voice,
speed=speed, speed=speed,
total_steps=total_steps,
split_pattern=split_pattern_override or params.tts_context.split_pattern, split_pattern=split_pattern_override or params.tts_context.split_pattern,
chapter_sink=chapter_sink, chapter_sink=chapter_sink,
preview_callback=preview_callback, preview_callback=preview_callback,
+9 -2
View File
@@ -146,6 +146,7 @@ def tts_segments(
speed: float, speed: float,
split_pattern: str, split_pattern: str,
current_time: float = 0.0, current_time: float = 0.0,
total_steps: Optional[int] = None,
) -> Iterator[SegmentResult]: ) -> Iterator[SegmentResult]:
"""Invoke TTS backend on (already normalized) text and yield SegmentResults. """Invoke TTS backend on (already normalized) text and yield SegmentResults.
@@ -159,16 +160,20 @@ def tts_segments(
speed: TTS speed multiplier. speed: TTS speed multiplier.
split_pattern: Regex pattern for sentence splitting. split_pattern: Regex pattern for sentence splitting.
current_time: Current position in the audio timeline (seconds). current_time: Current position in the audio timeline (seconds).
total_steps: Inference quality steps (Supertonic only, ignored by Kokoro).
Yields: Yields:
SegmentResult for each non-empty TTS segment. SegmentResult for each non-empty TTS segment.
""" """
segment_iter = backend( kwargs: dict[str, Any] = dict(
text,
voice=voice, voice=voice,
speed=speed, speed=speed,
split_pattern=split_pattern, split_pattern=split_pattern,
) )
if total_steps is not None:
kwargs["total_steps"] = total_steps
segment_iter = backend(text, **kwargs)
chunk_start = current_time chunk_start = current_time
@@ -215,6 +220,7 @@ def emit_text_segments(
speed: float, speed: float,
split_pattern: str, split_pattern: str,
current_time: float = 0.0, current_time: float = 0.0,
total_steps: Optional[int] = None,
# normalization # normalization
heteronym_rules: Any = None, heteronym_rules: Any = None,
pronunciation_rules: Any = None, pronunciation_rules: Any = None,
@@ -265,6 +271,7 @@ def emit_text_segments(
speed=speed, speed=speed,
split_pattern=split_pattern, split_pattern=split_pattern,
current_time=current_time, current_time=current_time,
total_steps=total_steps,
) )
+52 -4
View File
@@ -128,8 +128,8 @@ class InputFormat(str, Enum):
class Language(str, Enum): class Language(str, Enum):
"""TTS language code (ISO 639-1 with region where needed). """TTS language code (ISO 639-1 with region where needed).
Each engine (Kokoro, Supertonic) maps these to its own Each engine maps these to its own internal language identifiers.
internal language identifiers. Engines report which languages they support via ``supported_languages()``.
""" """
EN_US = "en-US" EN_US = "en-US"
EN_GB = "en-GB" EN_GB = "en-GB"
@@ -140,6 +140,30 @@ class Language(str, Enum):
JA = "ja" JA = "ja"
PT_BR = "pt-BR" PT_BR = "pt-BR"
ZH = "zh" ZH = "zh"
AR = "ar"
BG = "bg"
CS = "cs"
DA = "da"
DE = "de"
EL = "el"
ET = "et"
FI = "fi"
HR = "hr"
HU = "hu"
ID = "id"
KO = "ko"
LT = "lt"
LV = "lv"
NL = "nl"
PL = "pl"
RO = "ro"
RU = "ru"
SK = "sk"
SL = "sl"
SV = "sv"
TR = "tr"
UK = "uk"
VI = "vi"
@property @property
def display_name(self) -> str: def display_name(self) -> str:
@@ -154,13 +178,37 @@ class Language(str, Enum):
"ja": "Japanese", "ja": "Japanese",
"pt-BR": "Brazilian Portuguese", "pt-BR": "Brazilian Portuguese",
"zh": "Mandarin Chinese", "zh": "Mandarin Chinese",
"ar": "Arabic",
"bg": "Bulgarian",
"cs": "Czech",
"da": "Danish",
"de": "German",
"el": "Greek",
"et": "Estonian",
"fi": "Finnish",
"hr": "Croatian",
"hu": "Hungarian",
"id": "Indonesian",
"ko": "Korean",
"lt": "Lithuanian",
"lv": "Latvian",
"nl": "Dutch",
"pl": "Polish",
"ro": "Romanian",
"ru": "Russian",
"sk": "Slovak",
"sl": "Slovenian",
"sv": "Swedish",
"tr": "Turkish",
"uk": "Ukrainian",
"vi": "Vietnamese",
} }
return _names[self.value] return _names[self.value]
@property @property
def is_cjk(self) -> bool: def is_cjk(self) -> bool:
"""True for CJK languages (Chinese, Japanese).""" """True for CJK languages (Chinese, Japanese, Korean)."""
return self in (self.ZH, self.JA) return self in (self.ZH, self.JA, self.KO)
@property @property
def supports_subtitle_tokens(self) -> bool: def supports_subtitle_tokens(self) -> bool:
+1 -1
View File
@@ -45,7 +45,7 @@ def create_pipeline_for_job(
provider = "kokoro" provider = "kokoro"
if provider == "supertonic": if provider == "supertonic":
return create_pipeline("supertonic") return create_pipeline("supertonic", language=language)
device = resolve_device(use_gpu) device = resolve_device(use_gpu)
return create_pipeline("kokoro", language=language, device=device) return create_pipeline("kokoro", language=language, device=device)
+1 -1
View File
@@ -116,7 +116,7 @@ def generate_preview_audio(
if provider == "supertonic": if provider == "supertonic":
from abogen.tts_plugin.utils import create_pipeline from abogen.tts_plugin.utils import create_pipeline
pipeline = create_pipeline("supertonic") pipeline = create_pipeline("supertonic", language=language)
segments = pipeline( segments = pipeline(
normalized_text, normalized_text,
voice=voice_spec, voice=voice_spec,
+3 -2
View File
@@ -32,11 +32,12 @@ from abogen.tts_plugin.types import EngineConfig
from .engine import SuperTonicEngine from .engine import SuperTonicEngine
def _load_supertonic_pipeline() -> Any: def _load_supertonic_pipeline(language: Any = None) -> Any:
"""Lazy-load SuperTonic dependencies and create pipeline.""" """Lazy-load SuperTonic dependencies and create pipeline."""
from plugins.supertonic.pipeline import SupertonicPipeline from plugins.supertonic.pipeline import SupertonicPipeline
return SupertonicPipeline( return SupertonicPipeline(
language=language,
sample_rate=24000, sample_rate=24000,
auto_download=True, auto_download=True,
total_steps=5, total_steps=5,
@@ -128,7 +129,7 @@ def create_engine(
EngineError: On failure. Cleans up partially created resources. EngineError: On failure. Cleans up partially created resources.
""" """
try: try:
pipeline = _load_supertonic_pipeline() pipeline = _load_supertonic_pipeline(language=config.language)
engine = SuperTonicEngine(pipeline) engine = SuperTonicEngine(pipeline)
return engine return engine
except Exception as e: except Exception as e:
+56
View File
@@ -12,6 +12,7 @@ from typing import Any
import numpy as np import numpy as np
from abogen.domain.enums import Language
from abogen.tts_plugin.capabilities import VoiceLister from abogen.tts_plugin.capabilities import VoiceLister
from abogen.tts_plugin.engine import Engine, EngineSession from abogen.tts_plugin.engine import Engine, EngineSession
from abogen.tts_plugin.errors import EngineError from abogen.tts_plugin.errors import EngineError
@@ -28,6 +29,61 @@ logger = logging.getLogger(__name__)
# Sample rate for SuperTonic audio # Sample rate for SuperTonic audio
_SUPERTONIC_SAMPLE_RATE = 24000 _SUPERTONIC_SAMPLE_RATE = 24000
# Engine-internal language mapping: Language enum → Supertonic ISO 639-1 code.
_SUPERTONIC_LANG_MAP: dict[Language, str] = {
Language.EN_US: "en",
Language.EN_GB: "en",
Language.AR: "ar",
Language.BG: "bg",
Language.CS: "cs",
Language.DA: "da",
Language.DE: "de",
Language.EL: "el",
Language.ES: "es",
Language.ET: "et",
Language.FI: "fi",
Language.FR: "fr",
Language.HI: "hi",
Language.HR: "hr",
Language.HU: "hu",
Language.ID: "id",
Language.IT: "it",
Language.JA: "ja",
Language.KO: "ko",
Language.LT: "lt",
Language.LV: "lv",
Language.NL: "nl",
Language.PL: "pl",
Language.PT_BR: "pt",
Language.RO: "ro",
Language.RU: "ru",
Language.SK: "sk",
Language.SL: "sl",
Language.SV: "sv",
Language.TR: "tr",
Language.UK: "uk",
Language.VI: "vi",
}
def supported_languages() -> list[Language]:
"""Return the list of Language enum values this engine supports."""
return list(_SUPERTONIC_LANG_MAP.keys())
def engine_language(lang: Language) -> str:
"""Map a Language enum to the engine's internal ISO 639-1 code.
Raises ValueError for unsupported languages.
"""
result = _SUPERTONIC_LANG_MAP.get(lang)
if result is None:
raise ValueError(
f"Supertonic does not support language: {lang!r}. "
f"Supported: {supported_languages()}"
)
return result
class SuperTonicSession: class SuperTonicSession:
"""EngineSession implementation for SuperTonic. """EngineSession implementation for SuperTonic.
+9
View File
@@ -158,6 +158,7 @@ class SupertonicPipeline:
def __init__( def __init__(
self, self,
*, *,
language: Any = None,
sample_rate: int, sample_rate: int,
auto_download: bool = True, auto_download: bool = True,
total_steps: int = 5, total_steps: int = 5,
@@ -167,6 +168,13 @@ class SupertonicPipeline:
self.total_steps = int(total_steps) self.total_steps = int(total_steps)
self.max_chunk_length = int(max_chunk_length) self.max_chunk_length = int(max_chunk_length)
# Resolve language to ISO 639-1 code for Supertonic
if language is not None:
from plugins.supertonic.engine import engine_language
self._lang = engine_language(language)
else:
self._lang = "en"
_configure_supertonic_gpu() _configure_supertonic_gpu()
try: try:
@@ -212,6 +220,7 @@ class SupertonicPipeline:
max_chunk_length=self.max_chunk_length, max_chunk_length=self.max_chunk_length,
silence_duration=0.0, silence_duration=0.0,
verbose=False, verbose=False,
lang=self._lang,
) )
break break
except ValueError as exc: except ValueError as exc:
+1 -1
View File
@@ -54,7 +54,7 @@ class FakeBackend:
def __init__(self): def __init__(self):
self.synthesized: List[str] = [] self.synthesized: List[str] = []
def __call__(self, text: str, *, voice: Any, speed: float = 1.0, split_pattern: str = "") -> List: def __call__(self, text: str, *, voice: Any, speed: float = 1.0, split_pattern: str = "", **kwargs: Any) -> List:
self.synthesized.append(text) self.synthesized.append(text)
class FakeSegment: class FakeSegment:
+1 -1
View File
@@ -58,7 +58,7 @@ class FakeBackend:
self.segment_duration = segment_duration self.segment_duration = segment_duration
self.call_count = 0 self.call_count = 0
def __call__(self, text: str, voice: Any, speed: float = 1.0, split_pattern: str = ""): def __call__(self, text: str, voice: Any, speed: float = 1.0, split_pattern: str = "", **kwargs: Any):
self.call_count += 1 self.call_count += 1
# Return fake segment objects with required attributes # Return fake segment objects with required attributes
@dataclass @dataclass
+1 -1
View File
@@ -73,7 +73,7 @@ class FakeBackend:
def __init__(self): def __init__(self):
self.synthesized: List[str] = [] self.synthesized: List[str] = []
def __call__(self, text: str, *, voice: Any, speed: float = 1.0, split_pattern: str = "") -> List: def __call__(self, text: str, *, voice: Any, speed: float = 1.0, split_pattern: str = "", **kwargs: Any) -> List:
"""Return fake TTS segments.""" """Return fake TTS segments."""
self.synthesized.append(text) self.synthesized.append(text)
+1 -1
View File
@@ -33,7 +33,7 @@ class TestCreatePipelineForJob:
def test_supertonic_provider(self, _reg, mock_create): def test_supertonic_provider(self, _reg, mock_create):
mock_create.return_value = MagicMock() mock_create.return_value = MagicMock()
result = create_pipeline_for_job("supertonic", Language.EN_US, use_gpu=True) result = create_pipeline_for_job("supertonic", Language.EN_US, use_gpu=True)
mock_create.assert_called_once_with("supertonic") mock_create.assert_called_once_with("supertonic", language=Language.EN_US)
assert result is mock_create.return_value assert result is mock_create.return_value
@patch("abogen.domain.pipeline_factory.create_pipeline") @patch("abogen.domain.pipeline_factory.create_pipeline")
+1 -1
View File
@@ -49,7 +49,7 @@ def _make_mock_engine() -> Any:
from plugins.kokoro.engine import KokoroEngine from plugins.kokoro.engine import KokoroEngine
class MockPipeline: class MockPipeline:
def __call__(self, text, voice, speed, split_pattern=None): def __call__(self, text, voice, speed, split_pattern=None, **kwargs):
class MockSegment: class MockSegment:
def __init__(self): def __init__(self):
self.audio = MockAudio() self.audio = MockAudio()