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

@@ -25,6 +25,7 @@ from pipecat.services.openai import OpenAILLMService
from pipecat.services.playht import PlayHTTTSService from pipecat.services.playht import PlayHTTTSService
from pipecat.transports.services.daily import DailyParams, DailyTransport from pipecat.transports.services.daily import DailyParams, DailyTransport
from pipecat.vad.silero import SileroVADAnalyzer from pipecat.vad.silero import SileroVADAnalyzer
from pipecat.transcriptions.language import Language
load_dotenv(override=True) load_dotenv(override=True)
@@ -53,6 +54,7 @@ async def main():
user_id=os.getenv("PLAYHT_USER_ID"), user_id=os.getenv("PLAYHT_USER_ID"),
api_key=os.getenv("PLAYHT_API_KEY"), api_key=os.getenv("PLAYHT_API_KEY"),
voice_url="s3://voice-cloning-zero-shot/801a663f-efd0-4254-98d0-5c175514c3e8/jennifer/manifest.json", voice_url="s3://voice-cloning-zero-shot/801a663f-efd0-4254-98d0-5c175514c3e8/jennifer/manifest.json",
params=PlayHTTTSService.InputParams(language=Language.EN),
) )
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"), model="gpt-4o") llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"), model="gpt-4o")

View File

@@ -39,13 +39,13 @@ anthropic = [ "anthropic~=0.34.0" ]
aws = [ "boto3~=1.35.27" ] aws = [ "boto3~=1.35.27" ]
azure = [ "azure-cognitiveservices-speech~=1.40.0" ] azure = [ "azure-cognitiveservices-speech~=1.40.0" ]
canonical = [ "aiofiles~=24.1.0" ] canonical = [ "aiofiles~=24.1.0" ]
cartesia = [ "cartesia~=1.0.13", "websockets~=12.0" ] cartesia = [ "cartesia~=1.0.13", "websockets~=13.1" ]
daily = [ "daily-python~=0.11.0" ] daily = [ "daily-python~=0.11.0" ]
deepgram = [ "deepgram-sdk~=3.7.3" ] deepgram = [ "deepgram-sdk~=3.7.3" ]
elevenlabs = [ "websockets~=12.0" ] elevenlabs = [ "websockets~=13.1" ]
examples = [ "python-dotenv~=1.0.1", "flask~=3.0.3", "flask_cors~=4.0.1" ] examples = [ "python-dotenv~=1.0.1", "flask~=3.0.3", "flask_cors~=4.0.1" ]
fal = [ "fal-client~=0.4.1" ] fal = [ "fal-client~=0.4.1" ]
gladia = [ "websockets~=12.0" ] gladia = [ "websockets~=13.1" ]
google = [ "google-generativeai~=0.7.2", "google-cloud-texttospeech~=2.17.2" ] google = [ "google-generativeai~=0.7.2", "google-cloud-texttospeech~=2.17.2" ]
gstreamer = [ "pygobject~=3.48.2" ] gstreamer = [ "pygobject~=3.48.2" ]
fireworks = [ "openai~=1.37.2" ] fireworks = [ "openai~=1.37.2" ]
@@ -54,12 +54,12 @@ livekit = [ "livekit~=0.13.1", "tenacity~=9.0.0" ]
lmnt = [ "lmnt~=1.1.4" ] lmnt = [ "lmnt~=1.1.4" ]
local = [ "pyaudio~=0.2.14" ] local = [ "pyaudio~=0.2.14" ]
moondream = [ "einops~=0.8.0", "timm~=1.0.8", "transformers~=4.44.0" ] moondream = [ "einops~=0.8.0", "timm~=1.0.8", "transformers~=4.44.0" ]
openai = [ "openai~=1.50.2", "websockets~=12.0", "python-deepcompare~=1.0.1" ] openai = [ "openai~=1.50.2", "websockets~=13.1", "python-deepcompare~=1.0.1" ]
openpipe = [ "openpipe~=4.24.0" ] openpipe = [ "openpipe~=4.24.0" ]
playht = [ "pyht~=0.0.28" ] playht = [ "pyht~=0.1.4" ]
silero = [ "onnxruntime>=1.16.1" ] silero = [ "onnxruntime>=1.16.1" ]
together = [ "openai~=1.50.2" ] together = [ "openai~=1.50.2" ]
websocket = [ "websockets~=12.0", "fastapi~=0.115.0" ] websocket = [ "websockets~=13.1", "fastapi~=0.115.0" ]
whisper = [ "faster-whisper~=1.0.3" ] whisper = [ "faster-whisper~=1.0.3" ]
xtts = [ "resampy~=0.4.3" ] xtts = [ "resampy~=0.4.3" ]

View File

@@ -6,9 +6,10 @@
import io import io
import struct import struct
from typing import AsyncGenerator from typing import AsyncGenerator, Optional
from loguru import logger from loguru import logger
from pydantic.main import BaseModel
from pipecat.frames.frames import ( from pipecat.frames.frames import (
Frame, Frame,
@@ -17,6 +18,7 @@ from pipecat.frames.frames import (
TTSStoppedFrame, TTSStoppedFrame,
) )
from pipecat.services.ai_services import TTSService from pipecat.services.ai_services import TTSService
from pipecat.transcriptions.language import Language
try: try:
from pyht.async_client import AsyncClient from pyht.async_client import AsyncClient
@@ -31,8 +33,21 @@ except ModuleNotFoundError as e:
class PlayHTTTSService(TTSService): class PlayHTTTSService(TTSService):
class InputParams(BaseModel):
language: Optional[Language] = Language.EN
speed: Optional[float] = 1.0
seed: Optional[int] = None
def __init__( 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) super().__init__(sample_rate=sample_rate, **kwargs)
@@ -45,21 +60,87 @@ class PlayHTTTSService(TTSService):
) )
self._settings = { self._settings = {
"sample_rate": sample_rate, "sample_rate": sample_rate,
"quality": "higher", "language": self.language_to_service_language(params.language)
if params.language
else Language.EN,
"format": Format.FORMAT_WAV, "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.set_voice(voice_url)
self._options = TTSOptions( self._options = TTSOptions(
voice=self._voice_id, voice=self._voice_id,
language=self._settings["language"],
sample_rate=self._settings["sample_rate"], sample_rate=self._settings["sample_rate"],
quality=self._settings["quality"],
format=self._settings["format"], format=self._settings["format"],
speed=self._settings["speed"],
seed=self._settings["seed"],
) )
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 "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]: async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
logger.debug(f"Generating TTS: [{text}]") logger.debug(f"Generating TTS: [{text}]")

View File

@@ -20,10 +20,10 @@ Pillow~=10.4.0
pyaudio~=0.2.14 pyaudio~=0.2.14
pydantic~=2.8.2 pydantic~=2.8.2
pyloudnorm~=0.1.1 pyloudnorm~=0.1.1
pyht~=0.0.28 pyht~=0.1.4
python-dotenv~=1.0.1 python-dotenv~=1.0.1
resampy~=0.4.3 resampy~=0.4.3
silero-vad~=5.1 silero-vad~=5.1
together~=1.2.7 together~=1.2.7
transformers~=4.44.0 transformers~=4.44.0
websockets~=12.0 websockets~=13.1