services(riva): cleanup
This commit is contained in:
@@ -19,11 +19,8 @@ from pipecat.pipeline.pipeline import Pipeline
|
|||||||
from pipecat.pipeline.runner import PipelineRunner
|
from pipecat.pipeline.runner import PipelineRunner
|
||||||
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
from pipecat.pipeline.task import PipelineParams, PipelineTask
|
||||||
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext
|
from pipecat.processors.aggregators.openai_llm_context import OpenAILLMContext
|
||||||
|
from pipecat.services.nim import NimLLMService
|
||||||
from pipecat.services.deepgram import DeepgramSTTService
|
|
||||||
from pipecat.services.riva import FastpitchTTSService, ParakeetSTTService
|
from pipecat.services.riva import FastpitchTTSService, ParakeetSTTService
|
||||||
|
|
||||||
from pipecat.services.openai import OpenAILLMService
|
|
||||||
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
from pipecat.transports.services.daily import DailyParams, DailyTransport
|
||||||
|
|
||||||
load_dotenv(override=True)
|
load_dotenv(override=True)
|
||||||
@@ -48,12 +45,13 @@ async def main():
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# stt = DeepgramSTTService(api_key=os.getenv("DEEPGRAM_API_KEY"))
|
|
||||||
stt = ParakeetSTTService(api_key=os.getenv("NVIDIA_API_KEY"))
|
stt = ParakeetSTTService(api_key=os.getenv("NVIDIA_API_KEY"))
|
||||||
|
|
||||||
tts = FastpitchTTSService(api_key=os.getenv("NVIDIA_API_KEY"))
|
tts = FastpitchTTSService(api_key=os.getenv("NVIDIA_API_KEY"))
|
||||||
|
|
||||||
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"), model="gpt-4o")
|
llm = NimLLMService(
|
||||||
|
api_key=os.getenv("NVIDIA_API_KEY"), model="meta/llama-3.1-405b-instruct"
|
||||||
|
)
|
||||||
|
|
||||||
messages = [
|
messages = [
|
||||||
{
|
{
|
||||||
@@ -22,6 +22,7 @@ from pipecat.frames.frames import (
|
|||||||
TTSStoppedFrame,
|
TTSStoppedFrame,
|
||||||
)
|
)
|
||||||
from pipecat.services.ai_services import STTService, TTSService
|
from pipecat.services.ai_services import STTService, TTSService
|
||||||
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -37,7 +38,7 @@ except ModuleNotFoundError as e:
|
|||||||
|
|
||||||
class FastpitchTTSService(TTSService):
|
class FastpitchTTSService(TTSService):
|
||||||
class InputParams(BaseModel):
|
class InputParams(BaseModel):
|
||||||
language: Optional[str] = "en-US"
|
language: Optional[Language] = Language.EN_US
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -49,19 +50,19 @@ class FastpitchTTSService(TTSService):
|
|||||||
# nvidia riva calls this 'function-id'
|
# nvidia riva calls this 'function-id'
|
||||||
model: str = "0149dedb-2be8-4195-b9a0-e57e0e14f972",
|
model: str = "0149dedb-2be8-4195-b9a0-e57e0e14f972",
|
||||||
params: InputParams = InputParams(),
|
params: InputParams = InputParams(),
|
||||||
|
quality: int = 20,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(sample_rate=sample_rate_hz, **kwargs)
|
super().__init__(sample_rate=sample_rate_hz, **kwargs)
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
|
self._voice_id = voice_id
|
||||||
|
self._sample_rate_hz = sample_rate_hz
|
||||||
|
self._language_code = params.language
|
||||||
|
self.quality = quality
|
||||||
|
|
||||||
self.set_model_name("fastpitch-hifigan-tts")
|
self.set_model_name("fastpitch-hifigan-tts")
|
||||||
self.set_voice(voice_id)
|
self.set_voice(voice_id)
|
||||||
|
|
||||||
self.voice_id = voice_id
|
|
||||||
self.sample_rate_hz = sample_rate_hz
|
|
||||||
self.language_code = params.language
|
|
||||||
self.quality = None
|
|
||||||
|
|
||||||
metadata = [
|
metadata = [
|
||||||
["function-id", model],
|
["function-id", model],
|
||||||
["authorization", f"Bearer {api_key}"],
|
["authorization", f"Bearer {api_key}"],
|
||||||
@@ -73,15 +74,14 @@ class FastpitchTTSService(TTSService):
|
|||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
def read_audio_responses():
|
def read_audio_responses():
|
||||||
try:
|
try:
|
||||||
custom_dictionary_input = {}
|
|
||||||
responses = self.service.synthesize_online(
|
responses = self.service.synthesize_online(
|
||||||
text,
|
text,
|
||||||
self.voice_id,
|
self._voice_id,
|
||||||
self.language_code,
|
self._language_code,
|
||||||
sample_rate_hz=self.sample_rate_hz,
|
sample_rate_hz=self._sample_rate_hz,
|
||||||
audio_prompt_file=None,
|
audio_prompt_file=None,
|
||||||
quality=20 if self.quality is None else self.quality,
|
quality=self.quality,
|
||||||
custom_dictionary=custom_dictionary_input,
|
custom_dictionary={},
|
||||||
)
|
)
|
||||||
return responses
|
return responses
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -100,7 +100,7 @@ class FastpitchTTSService(TTSService):
|
|||||||
|
|
||||||
frame = TTSAudioRawFrame(
|
frame = TTSAudioRawFrame(
|
||||||
audio=resp.audio,
|
audio=resp.audio,
|
||||||
sample_rate=self.sample_rate_hz,
|
sample_rate=self._sample_rate_hz,
|
||||||
num_channels=1,
|
num_channels=1,
|
||||||
)
|
)
|
||||||
yield frame
|
yield frame
|
||||||
@@ -110,6 +110,9 @@ class FastpitchTTSService(TTSService):
|
|||||||
|
|
||||||
|
|
||||||
class ParakeetSTTService(STTService):
|
class ParakeetSTTService(STTService):
|
||||||
|
class InputParams(BaseModel):
|
||||||
|
language: Optional[Language] = Language.EN_US
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -117,28 +120,28 @@ class ParakeetSTTService(STTService):
|
|||||||
server: str = "grpc.nvcf.nvidia.com:443",
|
server: str = "grpc.nvcf.nvidia.com:443",
|
||||||
# nvidia calls this 'function-id'
|
# nvidia calls this 'function-id'
|
||||||
model: str = "1598d209-5e27-4d3c-8079-4751568b1081",
|
model: str = "1598d209-5e27-4d3c-8079-4751568b1081",
|
||||||
|
params: InputParams = InputParams(),
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
|
self._profanity_filter = False
|
||||||
|
self._automatic_punctuation = False
|
||||||
|
self._no_verbatim_transcripts = False
|
||||||
|
self._language_code = params.language
|
||||||
|
self._boosted_lm_words = None
|
||||||
|
self._boosted_lm_score = 4.0
|
||||||
|
self._start_history = -1
|
||||||
|
self._start_threshold = -1.0
|
||||||
|
self._stop_history = -1
|
||||||
|
self._stop_threshold = -1.0
|
||||||
|
self._stop_history_eou = -1
|
||||||
|
self._stop_threshold_eou = -1.0
|
||||||
|
self._custom_configuration = ""
|
||||||
|
self._sample_rate_hz: int = 16000
|
||||||
|
|
||||||
self.set_model_name("parakeet-ctc-1.1b-asr")
|
self.set_model_name("parakeet-ctc-1.1b-asr")
|
||||||
|
|
||||||
profanity_filter = False
|
|
||||||
automatic_punctuation = False
|
|
||||||
no_verbatim_transcripts = False
|
|
||||||
language_code = "en-US"
|
|
||||||
boosted_lm_words = None
|
|
||||||
boosted_lm_score = 4.0
|
|
||||||
start_history = -1
|
|
||||||
start_threshold = -1.0
|
|
||||||
stop_history = -1
|
|
||||||
stop_threshold = -1.0
|
|
||||||
stop_history_eou = -1
|
|
||||||
stop_threshold_eou = -1.0
|
|
||||||
custom_configuration = ""
|
|
||||||
sample_rate_hz: int = 16000
|
|
||||||
|
|
||||||
metadata = [
|
metadata = [
|
||||||
["function-id", model],
|
["function-id", model],
|
||||||
["authorization", f"Bearer {api_key}"],
|
["authorization", f"Bearer {api_key}"],
|
||||||
@@ -150,31 +153,32 @@ class ParakeetSTTService(STTService):
|
|||||||
config = riva.client.StreamingRecognitionConfig(
|
config = riva.client.StreamingRecognitionConfig(
|
||||||
config=riva.client.RecognitionConfig(
|
config=riva.client.RecognitionConfig(
|
||||||
encoding=riva.client.AudioEncoding.LINEAR_PCM,
|
encoding=riva.client.AudioEncoding.LINEAR_PCM,
|
||||||
language_code=language_code,
|
language_code=self._language_code,
|
||||||
model="",
|
model="",
|
||||||
max_alternatives=1,
|
max_alternatives=1,
|
||||||
profanity_filter=profanity_filter,
|
profanity_filter=self._profanity_filter,
|
||||||
enable_automatic_punctuation=automatic_punctuation,
|
enable_automatic_punctuation=self._automatic_punctuation,
|
||||||
verbatim_transcripts=not no_verbatim_transcripts,
|
verbatim_transcripts=not self._no_verbatim_transcripts,
|
||||||
sample_rate_hertz=sample_rate_hz,
|
sample_rate_hertz=self._sample_rate_hz,
|
||||||
audio_channel_count=1,
|
audio_channel_count=1,
|
||||||
),
|
),
|
||||||
interim_results=True,
|
interim_results=True,
|
||||||
)
|
)
|
||||||
self.config = config
|
riva.client.add_word_boosting_to_config(
|
||||||
riva.client.add_word_boosting_to_config(config, boosted_lm_words, boosted_lm_score)
|
config, self._boosted_lm_words, self._boosted_lm_score
|
||||||
|
)
|
||||||
riva.client.add_endpoint_parameters_to_config(
|
riva.client.add_endpoint_parameters_to_config(
|
||||||
config,
|
config,
|
||||||
start_history,
|
self._start_history,
|
||||||
start_threshold,
|
self._start_threshold,
|
||||||
stop_history,
|
self._stop_history,
|
||||||
stop_history_eou,
|
self._stop_history_eou,
|
||||||
stop_threshold,
|
self._stop_threshold,
|
||||||
stop_threshold_eou,
|
self._stop_threshold_eou,
|
||||||
)
|
)
|
||||||
riva.client.add_custom_configuration_to_config(config, custom_configuration)
|
riva.client.add_custom_configuration_to_config(config, self._custom_configuration)
|
||||||
|
self.config = config
|
||||||
|
|
||||||
# this doesn't work, but something like this perhaps? part 1
|
|
||||||
self._queue = asyncio.Queue()
|
self._queue = asyncio.Queue()
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
|
|||||||
Reference in New Issue
Block a user