mirror of
https://github.com/denizsafak/abogen.git
synced 2026-09-20 11:40:57 +02:00
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:
@@ -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]}"),
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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.
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
Reference in New Issue
Block a user