Merge pull request #545 from pipecat-ai/mb/fix-language-handling

Improve language string handling for TTS services
This commit is contained in:
Mark Backman
2024-10-04 10:03:06 -04:00
committed by GitHub
10 changed files with 505 additions and 485 deletions

View File

@@ -7,11 +7,16 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [0.0.43] - 2024-10-03 ## [0.0.43] - 2024-10-03
### Changed
- For TTS services, convert inputted languages to match each service's language
format
### Fixed ### Fixed
- Fixed an issue where changing a language with the Deepgram STT service - Fixed an issue where changing a language with the Deepgram STT service
wouldn't apply the change. This was fixed by disconnecting and reconnecting wouldn't apply the change. This was fixed by disconnecting and reconnecting
when the language changes when the language changes.
## [0.0.42] - 2024-10-02 ## [0.0.42] - 2024-10-02

View File

@@ -204,6 +204,9 @@ class TTSService(AIService):
async def flush_audio(self): async def flush_audio(self):
pass pass
def language_to_service_language(self, language: Language) -> str | None:
return Language(language)
# Converts the text to audio. # Converts the text to audio.
@abstractmethod @abstractmethod
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]: async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
@@ -234,7 +237,7 @@ class TTSService(AIService):
logger.debug(f"Updating TTS setting {key} to: [{value}]") logger.debug(f"Updating TTS setting {key} to: [{value}]")
self._settings[key] = value self._settings[key] = value
if key == "language": if key == "language":
self._settings[key] = Language(value) self._settings[key] = self.language_to_service_language(value)
elif key == "model": elif key == "model":
self.set_model_name(value) self.set_model_name(value)
elif key == "voice": elif key == "voice":

View File

@@ -30,67 +30,6 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
def language_to_aws_language(language: Language) -> str | None:
match language:
case Language.CA:
return "ca-ES"
case Language.ZH:
return "cmn-CN"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.NL_BE:
return "nl-BE"
case Language.EN:
return "en-US"
case Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_NZ:
return "en-NZ"
case Language.EN_IN:
return "en-IN"
case Language.FI:
return "fi-FI"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.HI:
return "hi-IN"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TR:
return "tr-TR"
return None
class AWSTTSService(TTSService): class AWSTTSService(TTSService):
class InputParams(BaseModel): class InputParams(BaseModel):
engine: Optional[str] = None engine: Optional[str] = None
@@ -121,7 +60,9 @@ class AWSTTSService(TTSService):
self._settings = { self._settings = {
"sample_rate": sample_rate, "sample_rate": sample_rate,
"engine": params.engine, "engine": params.engine,
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"pitch": params.pitch, "pitch": params.pitch,
"rate": params.rate, "rate": params.rate,
"volume": params.volume, "volume": params.volume,
@@ -132,10 +73,68 @@ class AWSTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.CA:
return "ca-ES"
case Language.ZH:
return "cmn-CN"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.NL_BE:
return "nl-BE"
case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_NZ:
return "en-NZ"
case Language.EN_IN:
return "en-IN"
case Language.FI:
return "fi-FI"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.HI:
return "hi-IN"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TR:
return "tr-TR"
return None
def _construct_ssml(self, text: str) -> str: def _construct_ssml(self, text: str) -> str:
ssml = "<speak>" ssml = "<speak>"
language = language_to_aws_language(self._settings["language"]) language = self._settings["language"]
ssml += f"<lang xml:lang='{language}'>" ssml += f"<lang xml:lang='{language}'>"
prosody_attrs = [] prosody_attrs = []

View File

@@ -27,7 +27,6 @@ from pipecat.frames.frames import (
) )
from pipecat.services.ai_services import ImageGenService, STTService, TTSService from pipecat.services.ai_services import ImageGenService, STTService, TTSService
from pipecat.services.openai import BaseOpenAILLMService from pipecat.services.openai import BaseOpenAILLMService
from pipecat.transcriptions import language
from pipecat.transcriptions.language import Language from pipecat.transcriptions.language import Language
from pipecat.utils.time import time_now_iso8601 from pipecat.utils.time import time_now_iso8601
@@ -72,101 +71,10 @@ class AzureLLMService(BaseOpenAILLMService):
) )
def language_to_azure_language(language: Language) -> str | None:
match language:
case Language.BG:
return "bg-BG"
case Language.CA:
return "ca-ES"
case Language.ZH:
return "zh-CN"
case Language.ZH_TW:
return "zh-TW"
case Language.CS:
return "cs-CZ"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.EN:
return "en-US"
case Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_NZ:
return "en-NZ"
case Language.EN_IN:
return "en-IN"
case Language.ET:
return "et-EE"
case Language.FI:
return "fi-FI"
case Language.NL_BE:
return "nl-BE"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.DE_CH:
return "de-CH"
case Language.EL:
return "el-GR"
case Language.HI:
return "hi-IN"
case Language.HU:
return "hu-HU"
case Language.ID:
return "id-ID"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.LV:
return "lv-LV"
case Language.LT:
return "lt-LT"
case Language.MS:
return "ms-MY"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.SK:
return "sk-SK"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TH:
return "th-TH"
case Language.TR:
return "tr-TR"
case Language.UK:
return "uk-UA"
case Language.VI:
return "vi-VN"
return None
class AzureTTSService(TTSService): class AzureTTSService(TTSService):
class InputParams(BaseModel): class InputParams(BaseModel):
emphasis: Optional[str] = None emphasis: Optional[str] = None
language: Optional[Language] = Language.EN language: Optional[Language] = Language.EN_US
pitch: Optional[str] = None pitch: Optional[str] = None
rate: Optional[str] = "1.05" rate: Optional[str] = "1.05"
role: Optional[str] = None role: Optional[str] = None
@@ -192,7 +100,9 @@ class AzureTTSService(TTSService):
self._settings = { self._settings = {
"sample_rate": sample_rate, "sample_rate": sample_rate,
"emphasis": params.emphasis, "emphasis": params.emphasis,
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN_US,
"pitch": params.pitch, "pitch": params.pitch,
"rate": params.rate, "rate": params.rate,
"role": params.role, "role": params.role,
@@ -206,8 +116,96 @@ class AzureTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.BG:
return "bg-BG"
case Language.CA:
return "ca-ES"
case Language.ZH:
return "zh-CN"
case Language.ZH_TW:
return "zh-TW"
case Language.CS:
return "cs-CZ"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_NZ:
return "en-NZ"
case Language.EN_IN:
return "en-IN"
case Language.ET:
return "et-EE"
case Language.FI:
return "fi-FI"
case Language.NL_BE:
return "nl-BE"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.DE_CH:
return "de-CH"
case Language.EL:
return "el-GR"
case Language.HI:
return "hi-IN"
case Language.HU:
return "hu-HU"
case Language.ID:
return "id-ID"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.LV:
return "lv-LV"
case Language.LT:
return "lt-LT"
case Language.MS:
return "ms-MY"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.SK:
return "sk-SK"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TH:
return "th-TH"
case Language.TR:
return "tr-TR"
case Language.UK:
return "uk-UA"
case Language.VI:
return "vi-VN"
return None
def _construct_ssml(self, text: str) -> str: def _construct_ssml(self, text: str) -> str:
language = language_to_azure_language(self._settings["language"]) language = self._settings["language"]
ssml = ( ssml = (
f"<speak version='1.0' xml:lang='{language}' " f"<speak version='1.0' xml:lang='{language}' "
"xmlns='http://www.w3.org/2001/10/synthesis' " "xmlns='http://www.w3.org/2001/10/synthesis' "
@@ -284,7 +282,7 @@ class AzureSTTService(STTService):
*, *,
api_key: str, api_key: str,
region: str, region: str,
language="en-US", language=Language.EN_US,
sample_rate=16000, sample_rate=16000,
channels=1, channels=1,
**kwargs, **kwargs,

View File

@@ -45,17 +45,24 @@ def language_to_cartesia_language(language: Language) -> str | None:
match language: match language:
case Language.DE: case Language.DE:
return "de" return "de"
case Language.EN: case (
Language.EN
| Language.EN_US
| Language.EN_GB
| Language.EN_AU
| Language.EN_NZ
| Language.EN_IN
):
return "en" return "en"
case Language.ES: case Language.ES:
return "es" return "es"
case Language.FR: case Language.FR | Language.FR_CA:
return "fr" return "fr"
case Language.JA: case Language.JA:
return "ja" return "ja"
case Language.PT: case Language.PT | Language.PT_BR:
return "pt" return "pt"
case Language.ZH: case Language.ZH | Language.ZH_TW:
return "zh" return "zh"
return None return None
@@ -106,7 +113,9 @@ class CartesiaTTSService(WordTTSService):
"encoding": params.encoding, "encoding": params.encoding,
"sample_rate": params.sample_rate, "sample_rate": params.sample_rate,
}, },
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"speed": params.speed, "speed": params.speed,
"emotion": params.emotion, "emotion": params.emotion,
} }
@@ -125,6 +134,9 @@ class CartesiaTTSService(WordTTSService):
await super().set_model(model) await super().set_model(model)
logger.debug(f"Switching TTS model to: [{model}]") logger.debug(f"Switching TTS model to: [{model}]")
def language_to_service_language(self, language: Language) -> str | None:
return language_to_cartesia_language(language)
def _build_msg( def _build_msg(
self, text: str = "", continue_transcript: bool = True, add_timestamps: bool = True self, text: str = "", continue_transcript: bool = True, add_timestamps: bool = True
): ):
@@ -146,7 +158,7 @@ class CartesiaTTSService(WordTTSService):
"model_id": self.model_name, "model_id": self.model_name,
"voice": voice_config, "voice": voice_config,
"output_format": self._settings["output_format"], "output_format": self._settings["output_format"],
"language": language_to_cartesia_language(self._settings["language"]), "language": self._settings["language"],
"add_timestamps": add_timestamps, "add_timestamps": add_timestamps,
} }
return json.dumps(msg) return json.dumps(msg)
@@ -303,7 +315,9 @@ class CartesiaHttpTTSService(TTSService):
"encoding": params.encoding, "encoding": params.encoding,
"sample_rate": params.sample_rate, "sample_rate": params.sample_rate,
}, },
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"speed": params.speed, "speed": params.speed,
"emotion": params.emotion, "emotion": params.emotion,
} }
@@ -315,6 +329,9 @@ class CartesiaHttpTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
return language_to_cartesia_language(language)
async def stop(self, frame: EndFrame): async def stop(self, frame: EndFrame):
await super().stop(frame) await super().stop(frame)
await self._client.close() await self._client.close()
@@ -343,7 +360,7 @@ class CartesiaHttpTTSService(TTSService):
transcript=text, transcript=text,
voice_id=self._voice_id, voice_id=self._voice_id,
output_format=self._settings["output_format"], output_format=self._settings["output_format"],
language=language_to_cartesia_language(self._settings["language"]), language=self._settings["language"],
stream=False, stream=False,
_experimental_voice_controls=voice_controls, _experimental_voice_controls=voice_controls,
) )

View File

@@ -50,76 +50,6 @@ def sample_rate_from_output_format(output_format: str) -> int:
return 16000 return 16000
def language_to_elevenlabs_language(language: Language) -> str | None:
match language:
case Language.BG:
return "bg"
case Language.ZH:
return "zh"
case Language.CS:
return "cs"
case Language.DA:
return "da"
case Language.NL:
return "nl"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.FI:
return "fi"
case Language.FR | Language.FR_CA:
return "fr"
case Language.DE | Language.DE_CH:
return "de"
case Language.EL:
return "el"
case Language.HI:
return "hi"
case Language.HU:
return "hu"
case Language.ID:
return "id"
case Language.IT:
return "it"
case Language.JA:
return "ja"
case Language.KO:
return "ko"
case Language.MS:
return "ms"
case Language.NO:
return "no"
case Language.PL:
return "pl"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro"
case Language.RU:
return "ru"
case Language.SK:
return "sk"
case Language.ES:
return "es"
case Language.SV:
return "sv"
case Language.TR:
return "tr"
case Language.UK:
return "uk"
case Language.VI:
return "vi"
return None
def calculate_word_times( def calculate_word_times(
alignment_info: Mapping[str, Any], cumulative_time: float alignment_info: Mapping[str, Any], cumulative_time: float
) -> List[Tuple[str, float]]: ) -> List[Tuple[str, float]]:
@@ -198,7 +128,9 @@ class ElevenLabsTTSService(WordTTSService):
self._url = url self._url = url
self._settings = { self._settings = {
"sample_rate": sample_rate_from_output_format(params.output_format), "sample_rate": sample_rate_from_output_format(params.output_format),
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"output_format": params.output_format, "output_format": params.output_format,
"optimize_streaming_latency": params.optimize_streaming_latency, "optimize_streaming_latency": params.optimize_streaming_latency,
"stability": params.stability, "stability": params.stability,
@@ -220,6 +152,75 @@ class ElevenLabsTTSService(WordTTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.BG:
return "bg"
case Language.ZH:
return "zh"
case Language.CS:
return "cs"
case Language.DA:
return "da"
case Language.NL:
return "nl"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.FI:
return "fi"
case Language.FR | Language.FR_CA:
return "fr"
case Language.DE | Language.DE_CH:
return "de"
case Language.EL:
return "el"
case Language.HI:
return "hi"
case Language.HU:
return "hu"
case Language.ID:
return "id"
case Language.IT:
return "it"
case Language.JA:
return "ja"
case Language.KO:
return "ko"
case Language.MS:
return "ms"
case Language.NO:
return "no"
case Language.PL:
return "pl"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro"
case Language.RU:
return "ru"
case Language.SK:
return "sk"
case Language.ES:
return "es"
case Language.SV:
return "sv"
case Language.TR:
return "tr"
case Language.UK:
return "uk"
case Language.VI:
return "vi"
return None
def _set_voice_settings(self): def _set_voice_settings(self):
voice_settings = {} voice_settings = {}
if ( if (
@@ -293,7 +294,7 @@ class ElevenLabsTTSService(WordTTSService):
url += f"&optimize_streaming_latency={self._settings['optimize_streaming_latency']}" url += f"&optimize_streaming_latency={self._settings['optimize_streaming_latency']}"
# Language can only be used with the 'eleven_turbo_v2_5' model # Language can only be used with the 'eleven_turbo_v2_5' model
language = language_to_elevenlabs_language(self._settings["language"]) language = self._settings["language"]
if model == "eleven_turbo_v2_5": if model == "eleven_turbo_v2_5":
url += f"&language_code={language}" url += f"&language_code={language}"
else: else:

View File

@@ -34,84 +34,6 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
def language_to_gladia_language(language: Language) -> str | None:
match language:
case Language.BG:
return "bulgarian"
case Language.CA:
return "catalan"
case Language.ZH:
return "chinese"
case Language.CS:
return "czech"
case Language.DA:
return "danish"
case Language.NL:
return "dutch"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "english"
case Language.ET:
return "estonian"
case Language.FI:
return "finnish"
case Language.FR | Language.FR_CA:
return "french"
case Language.DE | Language.DE_CH:
return "german"
case Language.EL:
return "greek"
case Language.HI:
return "hindi"
case Language.HU:
return "hungarian"
case Language.ID:
return "indonesian"
case Language.IT:
return "italian"
case Language.JA:
return "japanese"
case Language.KO:
return "korean"
case Language.LV:
return "latvian"
case Language.LT:
return "lithuanian"
case Language.MS:
return "malay"
case Language.NO:
return "norwegian"
case Language.PL:
return "polish"
case Language.PT | Language.PT_BR:
return "portuguese"
case Language.RO:
return "romanian"
case Language.RU:
return "russian"
case Language.SK:
return "slovak"
case Language.ES:
return "spanish"
case Language.SV:
return "slovenian"
case Language.TH:
return "thai"
case Language.TR:
return "turkish"
case Language.UK:
return "ukrainian"
case Language.VI:
return "vietnamese"
return None
class GladiaSTTService(STTService): class GladiaSTTService(STTService):
class InputParams(BaseModel): class InputParams(BaseModel):
sample_rate: Optional[int] = 16000 sample_rate: Optional[int] = 16000
@@ -135,13 +57,92 @@ class GladiaSTTService(STTService):
self._url = url self._url = url
self._settings = { self._settings = {
"sample_rate": params.sample_rate, "sample_rate": params.sample_rate,
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"transcription_hint": params.transcription_hint, "transcription_hint": params.transcription_hint,
"endpointing": params.endpointing, "endpointing": params.endpointing,
"prosody": params.prosody, "prosody": params.prosody,
} }
self._confidence = confidence self._confidence = confidence
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.BG:
return "bulgarian"
case Language.CA:
return "catalan"
case Language.ZH:
return "chinese"
case Language.CS:
return "czech"
case Language.DA:
return "danish"
case Language.NL:
return "dutch"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "english"
case Language.ET:
return "estonian"
case Language.FI:
return "finnish"
case Language.FR | Language.FR_CA:
return "french"
case Language.DE | Language.DE_CH:
return "german"
case Language.EL:
return "greek"
case Language.HI:
return "hindi"
case Language.HU:
return "hungarian"
case Language.ID:
return "indonesian"
case Language.IT:
return "italian"
case Language.JA:
return "japanese"
case Language.KO:
return "korean"
case Language.LV:
return "latvian"
case Language.LT:
return "lithuanian"
case Language.MS:
return "malay"
case Language.NO:
return "norwegian"
case Language.PL:
return "polish"
case Language.PT | Language.PT_BR:
return "portuguese"
case Language.RO:
return "romanian"
case Language.RU:
return "russian"
case Language.SK:
return "slovak"
case Language.ES:
return "spanish"
case Language.SV:
return "slovenian"
case Language.TH:
return "thai"
case Language.TR:
return "turkish"
case Language.UK:
return "ukrainian"
case Language.VI:
return "vietnamese"
return None
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
self._websocket = await websockets.connect(self._url) self._websocket = await websockets.connect(self._url)
@@ -169,7 +170,7 @@ class GladiaSTTService(STTService):
"model_type": "fast", "model_type": "fast",
"language_behaviour": "manual", "language_behaviour": "manual",
"sample_rate": self._settings["sample_rate"], "sample_rate": self._settings["sample_rate"],
"language": language_to_gladia_language(self._settings["language"]), "language": self._settings["language"],
"transcription_hint": self._settings["transcription_hint"], "transcription_hint": self._settings["transcription_hint"],
"endpointing": self._settings["endpointing"], "endpointing": self._settings["endpointing"],
"prosody": self._settings["prosody"], "prosody": self._settings["prosody"],

View File

@@ -146,93 +146,6 @@ class GoogleLLMService(LLMService):
await self._process_context(context) await self._process_context(context)
def language_to_google_language(language: Language) -> str | None:
match language:
case Language.BG:
return "bg-BG"
case Language.CA:
return "ca-ES"
case Language.ZH:
return "cmn-CN"
case Language.ZH_TW:
return "cmn-TW"
case Language.CS:
return "cs-CZ"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.EN:
return "en-US"
case Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_IN:
return "en-IN"
case Language.ET:
return "et-EE"
case Language.FI:
return "fi-FI"
case Language.NL_BE:
return "nl-BE"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.EL:
return "el-GR"
case Language.HI:
return "hi-IN"
case Language.HU:
return "hu-HU"
case Language.ID:
return "id-ID"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.LV:
return "lv-LV"
case Language.LT:
return "lt-LT"
case Language.MS:
return "ms-MY"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.SK:
return "sk-SK"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TH:
return "th-TH"
case Language.TR:
return "tr-TR"
case Language.UK:
return "uk-UA"
case Language.VI:
return "vi-VN"
return None
class GoogleTTSService(TTSService): class GoogleTTSService(TTSService):
class InputParams(BaseModel): class InputParams(BaseModel):
pitch: Optional[str] = None pitch: Optional[str] = None
@@ -261,7 +174,9 @@ class GoogleTTSService(TTSService):
"rate": params.rate, "rate": params.rate,
"volume": params.volume, "volume": params.volume,
"emphasis": params.emphasis, "emphasis": params.emphasis,
"language": params.language if params.language else Language.EN, "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"gender": params.gender, "gender": params.gender,
"google_style": params.google_style, "google_style": params.google_style,
} }
@@ -291,13 +206,97 @@ class GoogleTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.BG:
return "bg-BG"
case Language.CA:
return "ca-ES"
case Language.ZH:
return "cmn-CN"
case Language.ZH_TW:
return "cmn-TW"
case Language.CS:
return "cs-CZ"
case Language.DA:
return "da-DK"
case Language.NL:
return "nl-NL"
case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_AU:
return "en-AU"
case Language.EN_GB:
return "en-GB"
case Language.EN_IN:
return "en-IN"
case Language.ET:
return "et-EE"
case Language.FI:
return "fi-FI"
case Language.NL_BE:
return "nl-BE"
case Language.FR:
return "fr-FR"
case Language.FR_CA:
return "fr-CA"
case Language.DE:
return "de-DE"
case Language.EL:
return "el-GR"
case Language.HI:
return "hi-IN"
case Language.HU:
return "hu-HU"
case Language.ID:
return "id-ID"
case Language.IT:
return "it-IT"
case Language.JA:
return "ja-JP"
case Language.KO:
return "ko-KR"
case Language.LV:
return "lv-LV"
case Language.LT:
return "lt-LT"
case Language.MS:
return "ms-MY"
case Language.NO:
return "nb-NO"
case Language.PL:
return "pl-PL"
case Language.PT:
return "pt-PT"
case Language.PT_BR:
return "pt-BR"
case Language.RO:
return "ro-RO"
case Language.RU:
return "ru-RU"
case Language.SK:
return "sk-SK"
case Language.ES:
return "es-ES"
case Language.SV:
return "sv-SE"
case Language.TH:
return "th-TH"
case Language.TR:
return "tr-TR"
case Language.UK:
return "uk-UA"
case Language.VI:
return "vi-VN"
return None
def _construct_ssml(self, text: str) -> str: def _construct_ssml(self, text: str) -> str:
ssml = "<speak>" ssml = "<speak>"
# Voice tag # Voice tag
voice_attrs = [f"name='{self._voice_id}'"] voice_attrs = [f"name='{self._voice_id}'"]
language = language_to_google_language(self._settings["language"]) language = self._settings["language"]
voice_attrs.append(f"language='{language}'") voice_attrs.append(f"language='{language}'")
if self._settings["gender"]: if self._settings["gender"]:

View File

@@ -35,32 +35,6 @@ except ModuleNotFoundError as e:
raise Exception(f"Missing module: {e}") raise Exception(f"Missing module: {e}")
def language_to_lmnt_language(language: Language) -> str | None:
match language:
case Language.DE:
return "de"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.ES:
return "es"
case Language.FR | Language.FR_CA:
return "fr"
case Language.PT | Language.PT_BR:
return "pt"
case Language.ZH | Language.ZH_TW:
return "zh"
case Language.KO:
return "ko"
return None
class LmntTTSService(TTSService): class LmntTTSService(TTSService):
def __init__( def __init__(
self, self,
@@ -82,7 +56,7 @@ class LmntTTSService(TTSService):
"encoding": "pcm_s16le", "encoding": "pcm_s16le",
"sample_rate": sample_rate, "sample_rate": sample_rate,
}, },
"language": language, "language": self.language_to_service_language(language),
} }
self.set_voice(voice_id) self.set_voice(voice_id)
@@ -97,6 +71,31 @@ class LmntTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.DE:
return "de"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.ES:
return "es"
case Language.FR | Language.FR_CA:
return "fr"
case Language.PT | Language.PT_BR:
return "pt"
case Language.ZH | Language.ZH_TW:
return "zh"
case Language.KO:
return "ko"
return None
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
await self._connect() await self._connect()
@@ -121,6 +120,7 @@ class LmntTTSService(TTSService):
self._voice_id, self._voice_id,
format="raw", format="raw",
sample_rate=self._settings["output_format"]["sample_rate"], sample_rate=self._settings["output_format"]["sample_rate"],
language=self._settings["language"],
) )
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler()) self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
except Exception as e: except Exception as e:

View File

@@ -37,50 +37,6 @@ except ModuleNotFoundError as e:
# https://github.com/coqui-ai/xtts-streaming-server # https://github.com/coqui-ai/xtts-streaming-server
def language_to_xtts_language(language: Language) -> str | None:
match language:
case Language.CS:
return "cs"
case Language.DE:
return "de"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.ES:
return "es"
case Language.FR:
return "fr"
case Language.HI:
return "hi"
case Language.HU:
return "hu"
case Language.IT:
return "it"
case Language.JA:
return "ja"
case Language.KO:
return "ko"
case Language.NL:
return "nl"
case Language.PL:
return "pl"
case Language.PT | Language.PT_BR:
return "pt"
case Language.RU:
return "ru"
case Language.TR:
return "tr"
case Language.ZH:
return "zh-cn"
return None
class XTTSService(TTSService): class XTTSService(TTSService):
def __init__( def __init__(
self, self,
@@ -94,7 +50,7 @@ class XTTSService(TTSService):
super().__init__(**kwargs) super().__init__(**kwargs)
self._settings = { self._settings = {
"language": language, "language": self.language_to_service_language(language),
"base_url": base_url, "base_url": base_url,
} }
self.set_voice(voice_id) self.set_voice(voice_id)
@@ -104,6 +60,49 @@ class XTTSService(TTSService):
def can_generate_metrics(self) -> bool: def can_generate_metrics(self) -> bool:
return True return True
def language_to_service_language(self, language: Language) -> str | None:
match language:
case Language.CS:
return "cs"
case Language.DE:
return "de"
case (
Language.EN
| Language.EN_US
| Language.EN_AU
| Language.EN_GB
| Language.EN_NZ
| Language.EN_IN
):
return "en"
case Language.ES:
return "es"
case Language.FR:
return "fr"
case Language.HI:
return "hi"
case Language.HU:
return "hu"
case Language.IT:
return "it"
case Language.JA:
return "ja"
case Language.KO:
return "ko"
case Language.NL:
return "nl"
case Language.PL:
return "pl"
case Language.PT | Language.PT_BR:
return "pt"
case Language.RU:
return "ru"
case Language.TR:
return "tr"
case Language.ZH:
return "zh-cn"
return None
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
async with self._aiohttp_session.get(self._settings["base_url"] + "/studio_speakers") as r: async with self._aiohttp_session.get(self._settings["base_url"] + "/studio_speakers") as r:
@@ -131,11 +130,9 @@ class XTTSService(TTSService):
url = self._settings["base_url"] + "/tts_stream" url = self._settings["base_url"] + "/tts_stream"
language = language_to_xtts_language(self._settings["language"])
payload = { payload = {
"text": text.replace(".", "").replace("*", ""), "text": text.replace(".", "").replace("*", ""),
"language": language, "language": self._settings["language"],
"speaker_embedding": embeddings["speaker_embedding"], "speaker_embedding": embeddings["speaker_embedding"],
"gpt_cond_latent": embeddings["gpt_cond_latent"], "gpt_cond_latent": embeddings["gpt_cond_latent"],
"add_wav_header": False, "add_wav_header": False,