Add input options to PlayHT, upgrade to latest PlayHT model

This commit is contained in:
Mark Backman
2024-10-17 11:09:18 -04:00
parent e31d1152db
commit 45606e177c
4 changed files with 96 additions and 13 deletions

View File

@@ -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}]")