Add input options to PlayHT, upgrade to latest PlayHT model
This commit is contained in:
@@ -6,9 +6,10 @@
|
||||
|
||||
import io
|
||||
import struct
|
||||
from typing import AsyncGenerator
|
||||
from typing import AsyncGenerator, Optional
|
||||
|
||||
from loguru import logger
|
||||
from pydantic.main import BaseModel
|
||||
|
||||
from pipecat.frames.frames import (
|
||||
Frame,
|
||||
@@ -17,6 +18,7 @@ from pipecat.frames.frames import (
|
||||
TTSStoppedFrame,
|
||||
)
|
||||
from pipecat.services.ai_services import TTSService
|
||||
from pipecat.transcriptions.language import Language
|
||||
|
||||
try:
|
||||
from pyht.async_client import AsyncClient
|
||||
@@ -31,8 +33,21 @@ except ModuleNotFoundError as e:
|
||||
|
||||
|
||||
class PlayHTTTSService(TTSService):
|
||||
class InputParams(BaseModel):
|
||||
language: Optional[Language] = Language.EN
|
||||
speed: Optional[float] = 1.0
|
||||
seed: Optional[int] = None
|
||||
|
||||
def __init__(
|
||||
self, *, api_key: str, user_id: str, voice_url: str, sample_rate: int = 16000, **kwargs
|
||||
self,
|
||||
*,
|
||||
api_key: str,
|
||||
user_id: str,
|
||||
voice_url: str,
|
||||
voice_engine: str = "PlayHT3.0-mini",
|
||||
sample_rate: int = 16000,
|
||||
params: InputParams = InputParams(),
|
||||
**kwargs,
|
||||
):
|
||||
super().__init__(sample_rate=sample_rate, **kwargs)
|
||||
|
||||
@@ -45,21 +60,87 @@ class PlayHTTTSService(TTSService):
|
||||
)
|
||||
self._settings = {
|
||||
"sample_rate": sample_rate,
|
||||
"quality": "higher",
|
||||
"language": self.language_to_service_language(params.language)
|
||||
if params.language
|
||||
else Language.EN,
|
||||
"format": Format.FORMAT_WAV,
|
||||
"voice_engine": "PlayHT2.0-turbo",
|
||||
"voice_engine": voice_engine,
|
||||
"speed": params.speed,
|
||||
"seed": params.seed,
|
||||
}
|
||||
self.set_model_name(voice_engine)
|
||||
self.set_voice(voice_url)
|
||||
self._options = TTSOptions(
|
||||
voice=self._voice_id,
|
||||
language=self._settings["language"],
|
||||
sample_rate=self._settings["sample_rate"],
|
||||
quality=self._settings["quality"],
|
||||
format=self._settings["format"],
|
||||
speed=self._settings["speed"],
|
||||
seed=self._settings["seed"],
|
||||
)
|
||||
|
||||
def can_generate_metrics(self) -> bool:
|
||||
return True
|
||||
|
||||
def language_to_service_language(self, language: Language) -> str | None:
|
||||
match language:
|
||||
case Language.BG:
|
||||
return "BULGARIAN"
|
||||
case Language.CA:
|
||||
return "CATALAN"
|
||||
case Language.CS:
|
||||
return "CZECH"
|
||||
case Language.DA:
|
||||
return "DANISH"
|
||||
case Language.DE:
|
||||
return "GERMAN"
|
||||
case (
|
||||
Language.EN
|
||||
| Language.EN_US
|
||||
| Language.EN_GB
|
||||
| Language.EN_AU
|
||||
| Language.EN_NZ
|
||||
| Language.EN_IN
|
||||
):
|
||||
return "ENGLISH"
|
||||
case Language.ES:
|
||||
return "SPANISH"
|
||||
case Language.FR | Language.FR_CA:
|
||||
return "FRENCH"
|
||||
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.MS:
|
||||
return "MALAY"
|
||||
case Language.NL:
|
||||
return "DUTCH"
|
||||
case Language.PL:
|
||||
return "POLISH"
|
||||
case Language.PT | Language.PT_BR:
|
||||
return "PORTUGUESE"
|
||||
case Language.RU:
|
||||
return "RUSSIAN"
|
||||
case Language.SV:
|
||||
return "SWEDISH"
|
||||
case Language.TH:
|
||||
return "THAI"
|
||||
case Language.TR:
|
||||
return "TURKISH"
|
||||
case Language.UK:
|
||||
return "UKRAINIAN"
|
||||
return None
|
||||
|
||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||
logger.debug(f"Generating TTS: [{text}]")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user