Improve language handling for TTS services

This commit is contained in:
Mark Backman
2024-10-03 22:51:15 -04:00
parent 27dcf83f37
commit 7796a272ce
9 changed files with 499 additions and 484 deletions

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,7 +30,50 @@ 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: class AWSTTSService(TTSService):
class InputParams(BaseModel):
engine: Optional[str] = None
language: Optional[Language] = Language.EN
pitch: Optional[str] = None
rate: Optional[str] = None
volume: Optional[str] = None
def __init__(
self,
*,
api_key: str,
aws_access_key_id: str,
region: str,
voice_id: str = "Joanna",
sample_rate: int = 16000,
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(sample_rate=sample_rate, **kwargs)
self._polly_client = boto3.client(
"polly",
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=api_key,
region_name=region,
)
self._settings = {
"sample_rate": sample_rate,
"engine": params.engine,
"language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"pitch": params.pitch,
"rate": params.rate,
"volume": params.volume,
}
self.set_voice(voice_id)
def can_generate_metrics(self) -> bool:
return True
def language_to_service_language(self, language: Language) -> str | None:
match language: match language:
case Language.CA: case Language.CA:
return "ca-ES" return "ca-ES"
@@ -42,9 +85,7 @@ def language_to_aws_language(language: Language) -> str | None:
return "nl-NL" return "nl-NL"
case Language.NL_BE: case Language.NL_BE:
return "nl-BE" return "nl-BE"
case Language.EN: case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_US:
return "en-US" return "en-US"
case Language.EN_AU: case Language.EN_AU:
return "en-AU" return "en-AU"
@@ -90,52 +131,10 @@ def language_to_aws_language(language: Language) -> str | None:
return "tr-TR" return "tr-TR"
return None return None
class AWSTTSService(TTSService):
class InputParams(BaseModel):
engine: Optional[str] = None
language: Optional[Language] = Language.EN
pitch: Optional[str] = None
rate: Optional[str] = None
volume: Optional[str] = None
def __init__(
self,
*,
api_key: str,
aws_access_key_id: str,
region: str,
voice_id: str = "Joanna",
sample_rate: int = 16000,
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(sample_rate=sample_rate, **kwargs)
self._polly_client = boto3.client(
"polly",
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=api_key,
region_name=region,
)
self._settings = {
"sample_rate": sample_rate,
"engine": params.engine,
"language": params.language if params.language else Language.EN,
"pitch": params.pitch,
"rate": params.rate,
"volume": params.volume,
}
self.set_voice(voice_id)
def can_generate_metrics(self) -> bool:
return True
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,7 +71,52 @@ class AzureLLMService(BaseOpenAILLMService):
) )
def language_to_azure_language(language: Language) -> str | None: class AzureTTSService(TTSService):
class InputParams(BaseModel):
emphasis: Optional[str] = None
language: Optional[Language] = Language.EN_US
pitch: Optional[str] = None
rate: Optional[str] = "1.05"
role: Optional[str] = None
style: Optional[str] = None
style_degree: Optional[str] = None
volume: Optional[str] = None
def __init__(
self,
*,
api_key: str,
region: str,
voice="en-US-SaraNeural",
sample_rate: int = 16000,
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(sample_rate=sample_rate, **kwargs)
speech_config = SpeechConfig(subscription=api_key, region=region)
self._speech_synthesizer = SpeechSynthesizer(speech_config=speech_config, audio_config=None)
self._settings = {
"sample_rate": sample_rate,
"emphasis": params.emphasis,
"language": self.language_to_service_language(params.language)
if params.language
else Language.EN_US,
"pitch": params.pitch,
"rate": params.rate,
"role": params.role,
"style": params.style,
"style_degree": params.style_degree,
"volume": params.volume,
}
self.set_voice(voice)
def can_generate_metrics(self) -> bool:
return True
def language_to_service_language(self, language: Language) -> str | None:
match language: match language:
case Language.BG: case Language.BG:
return "bg-BG" return "bg-BG"
@@ -88,9 +132,7 @@ def language_to_azure_language(language: Language) -> str | None:
return "da-DK" return "da-DK"
case Language.NL: case Language.NL:
return "nl-NL" return "nl-NL"
case Language.EN: case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_US:
return "en-US" return "en-US"
case Language.EN_AU: case Language.EN_AU:
return "en-AU" return "en-AU"
@@ -162,52 +204,8 @@ def language_to_azure_language(language: Language) -> str | None:
return "vi-VN" return "vi-VN"
return None return None
class AzureTTSService(TTSService):
class InputParams(BaseModel):
emphasis: Optional[str] = None
language: Optional[Language] = Language.EN
pitch: Optional[str] = None
rate: Optional[str] = "1.05"
role: Optional[str] = None
style: Optional[str] = None
style_degree: Optional[str] = None
volume: Optional[str] = None
def __init__(
self,
*,
api_key: str,
region: str,
voice="en-US-SaraNeural",
sample_rate: int = 16000,
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(sample_rate=sample_rate, **kwargs)
speech_config = SpeechConfig(subscription=api_key, region=region)
self._speech_synthesizer = SpeechSynthesizer(speech_config=speech_config, audio_config=None)
self._settings = {
"sample_rate": sample_rate,
"emphasis": params.emphasis,
"language": params.language if params.language else Language.EN,
"pitch": params.pitch,
"rate": params.rate,
"role": params.role,
"style": params.style,
"style_degree": params.style_degree,
"volume": params.volume,
}
self.set_voice(voice)
def can_generate_metrics(self) -> bool:
return True
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,7 +34,39 @@ 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: class GladiaSTTService(STTService):
class InputParams(BaseModel):
sample_rate: Optional[int] = 16000
language: Optional[Language] = Language.EN
transcription_hint: Optional[str] = None
endpointing: Optional[int] = 200
prosody: Optional[bool] = None
def __init__(
self,
*,
api_key: str,
url: str = "wss://api.gladia.io/audio/text/audio-transcription",
confidence: float = 0.5,
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(**kwargs)
self._api_key = api_key
self._url = url
self._settings = {
"sample_rate": params.sample_rate,
"language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"transcription_hint": params.transcription_hint,
"endpointing": params.endpointing,
"prosody": params.prosody,
}
self._confidence = confidence
def language_to_service_language(self, language: Language) -> str | None:
match language: match language:
case Language.BG: case Language.BG:
return "bulgarian" return "bulgarian"
@@ -111,37 +143,6 @@ def language_to_gladia_language(language: Language) -> str | None:
return "vietnamese" return "vietnamese"
return None return None
class GladiaSTTService(STTService):
class InputParams(BaseModel):
sample_rate: Optional[int] = 16000
language: Optional[Language] = Language.EN
transcription_hint: Optional[str] = None
endpointing: Optional[int] = 200
prosody: Optional[bool] = None
def __init__(
self,
*,
api_key: str,
url: str = "wss://api.gladia.io/audio/text/audio-transcription",
confidence: float = 0.5,
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(**kwargs)
self._api_key = api_key
self._url = url
self._settings = {
"sample_rate": params.sample_rate,
"language": params.language if params.language else Language.EN,
"transcription_hint": params.transcription_hint,
"endpointing": params.endpointing,
"prosody": params.prosody,
}
self._confidence = confidence
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,7 +146,67 @@ class GoogleLLMService(LLMService):
await self._process_context(context) await self._process_context(context)
def language_to_google_language(language: Language) -> str | None: class GoogleTTSService(TTSService):
class InputParams(BaseModel):
pitch: Optional[str] = None
rate: Optional[str] = None
volume: Optional[str] = None
emphasis: Optional[Literal["strong", "moderate", "reduced", "none"]] = None
language: Optional[Language] = Language.EN
gender: Optional[Literal["male", "female", "neutral"]] = None
google_style: Optional[Literal["apologetic", "calm", "empathetic", "firm", "lively"]] = None
def __init__(
self,
*,
credentials: Optional[str] = None,
credentials_path: Optional[str] = None,
voice_id: str = "en-US-Neural2-A",
sample_rate: int = 24000,
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(sample_rate=sample_rate, **kwargs)
self._settings = {
"sample_rate": sample_rate,
"pitch": params.pitch,
"rate": params.rate,
"volume": params.volume,
"emphasis": params.emphasis,
"language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"gender": params.gender,
"google_style": params.google_style,
}
self.set_voice(voice_id)
self._client: texttospeech_v1.TextToSpeechAsyncClient = self._create_client(
credentials, credentials_path
)
def _create_client(
self, credentials: Optional[str], credentials_path: Optional[str]
) -> texttospeech_v1.TextToSpeechAsyncClient:
creds: Optional[service_account.Credentials] = None
# Create a Google Cloud service account for the Cloud Text-to-Speech API
# Using either the provided credentials JSON string or the path to a service account JSON
# file, create a Google Cloud service account and use it to authenticate with the API.
if credentials:
# Use provided credentials JSON string
json_account_info = json.loads(credentials)
creds = service_account.Credentials.from_service_account_info(json_account_info)
elif credentials_path:
# Use service account JSON file if provided
creds = service_account.Credentials.from_service_account_file(credentials_path)
return texttospeech_v1.TextToSpeechAsyncClient(credentials=creds)
def can_generate_metrics(self) -> bool:
return True
def language_to_service_language(self, language: Language) -> str | None:
match language: match language:
case Language.BG: case Language.BG:
return "bg-BG" return "bg-BG"
@@ -162,9 +222,7 @@ def language_to_google_language(language: Language) -> str | None:
return "da-DK" return "da-DK"
case Language.NL: case Language.NL:
return "nl-NL" return "nl-NL"
case Language.EN: case Language.EN | Language.EN_US:
return "en-US"
case Language.EN_US:
return "en-US" return "en-US"
case Language.EN_AU: case Language.EN_AU:
return "en-AU" return "en-AU"
@@ -232,72 +290,13 @@ def language_to_google_language(language: Language) -> str | None:
return "vi-VN" return "vi-VN"
return None return None
class GoogleTTSService(TTSService):
class InputParams(BaseModel):
pitch: Optional[str] = None
rate: Optional[str] = None
volume: Optional[str] = None
emphasis: Optional[Literal["strong", "moderate", "reduced", "none"]] = None
language: Optional[Language] = Language.EN
gender: Optional[Literal["male", "female", "neutral"]] = None
google_style: Optional[Literal["apologetic", "calm", "empathetic", "firm", "lively"]] = None
def __init__(
self,
*,
credentials: Optional[str] = None,
credentials_path: Optional[str] = None,
voice_id: str = "en-US-Neural2-A",
sample_rate: int = 24000,
params: InputParams = InputParams(),
**kwargs,
):
super().__init__(sample_rate=sample_rate, **kwargs)
self._settings = {
"sample_rate": sample_rate,
"pitch": params.pitch,
"rate": params.rate,
"volume": params.volume,
"emphasis": params.emphasis,
"language": params.language if params.language else Language.EN,
"gender": params.gender,
"google_style": params.google_style,
}
self.set_voice(voice_id)
self._client: texttospeech_v1.TextToSpeechAsyncClient = self._create_client(
credentials, credentials_path
)
def _create_client(
self, credentials: Optional[str], credentials_path: Optional[str]
) -> texttospeech_v1.TextToSpeechAsyncClient:
creds: Optional[service_account.Credentials] = None
# Create a Google Cloud service account for the Cloud Text-to-Speech API
# Using either the provided credentials JSON string or the path to a service account JSON
# file, create a Google Cloud service account and use it to authenticate with the API.
if credentials:
# Use provided credentials JSON string
json_account_info = json.loads(credentials)
creds = service_account.Credentials.from_service_account_info(json_account_info)
elif credentials_path:
# Use service account JSON file if provided
creds = service_account.Credentials.from_service_account_file(credentials_path)
return texttospeech_v1.TextToSpeechAsyncClient(credentials=creds)
def can_generate_metrics(self) -> bool:
return True
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,7 +35,43 @@ 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: class LmntTTSService(TTSService):
def __init__(
self,
*,
api_key: str,
voice_id: str,
sample_rate: int = 24000,
language: Language = Language.EN,
**kwargs,
):
# Let TTSService produce TTSStoppedFrames after a short delay of
# no activity.
super().__init__(push_stop_frames=True, sample_rate=sample_rate, **kwargs)
self._api_key = api_key
self._settings = {
"output_format": {
"container": "raw",
"encoding": "pcm_s16le",
"sample_rate": sample_rate,
},
"language": self.language_to_service_language(language),
}
self.set_voice(voice_id)
self._speech = None
self._connection = None
self._receive_task = None
# Indicates if we have sent TTSStartedFrame. It will reset to False when
# there's an interruption or TTSStoppedFrame.
self._started = False
def can_generate_metrics(self) -> bool:
return True
def language_to_service_language(self, language: Language) -> str | None:
match language: match language:
case Language.DE: case Language.DE:
return "de" return "de"
@@ -60,43 +96,6 @@ def language_to_lmnt_language(language: Language) -> str | None:
return "ko" return "ko"
return None return None
class LmntTTSService(TTSService):
def __init__(
self,
*,
api_key: str,
voice_id: str,
sample_rate: int = 24000,
language: Language = Language.EN,
**kwargs,
):
# Let TTSService produce TTSStoppedFrames after a short delay of
# no activity.
super().__init__(push_stop_frames=True, sample_rate=sample_rate, **kwargs)
self._api_key = api_key
self._settings = {
"output_format": {
"container": "raw",
"encoding": "pcm_s16le",
"sample_rate": sample_rate,
},
"language": language,
}
self.set_voice(voice_id)
self._speech = None
self._connection = None
self._receive_task = None
# Indicates if we have sent TTSStartedFrame. It will reset to False when
# there's an interruption or TTSStoppedFrame.
self._started = False
def can_generate_metrics(self) -> bool:
return True
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,7 +37,30 @@ 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: class XTTSService(TTSService):
def __init__(
self,
*,
voice_id: str,
language: Language,
base_url: str,
aiohttp_session: aiohttp.ClientSession,
**kwargs,
):
super().__init__(**kwargs)
self._settings = {
"language": self.language_to_service_language(language),
"base_url": base_url,
}
self.set_voice(voice_id)
self._studio_speakers: Dict[str, Any] | None = None
self._aiohttp_session = aiohttp_session
def can_generate_metrics(self) -> bool:
return True
def language_to_service_language(self, language: Language) -> str | None:
match language: match language:
case Language.CS: case Language.CS:
return "cs" return "cs"
@@ -80,30 +103,6 @@ def language_to_xtts_language(language: Language) -> str | None:
return "zh-cn" return "zh-cn"
return None return None
class XTTSService(TTSService):
def __init__(
self,
*,
voice_id: str,
language: Language,
base_url: str,
aiohttp_session: aiohttp.ClientSession,
**kwargs,
):
super().__init__(**kwargs)
self._settings = {
"language": language,
"base_url": base_url,
}
self.set_voice(voice_id)
self._studio_speakers: Dict[str, Any] | None = None
self._aiohttp_session = aiohttp_session
def can_generate_metrics(self) -> bool:
return True
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,