Merge pull request #3509 from pipecat-ai/aleix/nvidia-stt-tts-improvements
NVIDIA STT/TTS performance improvements
This commit is contained in:
1
changelog/3509.fixed.2.md
Normal file
1
changelog/3509.fixed.2.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Optimized `NVIDIATTSService` to process incoming audio frames immediately.
|
||||||
1
changelog/3509.fixed.md
Normal file
1
changelog/3509.fixed.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Optimized `NVIDIASTTService` by removing unnecessary queue and task.
|
||||||
@@ -167,7 +167,6 @@ class NvidiaSTTService(STTService):
|
|||||||
self._queue = None
|
self._queue = None
|
||||||
self._config = None
|
self._config = None
|
||||||
self._thread_task = None
|
self._thread_task = None
|
||||||
self._response_task = None
|
|
||||||
|
|
||||||
def _initialize_client(self):
|
def _initialize_client(self):
|
||||||
metadata = [
|
metadata = [
|
||||||
@@ -251,10 +250,6 @@ class NvidiaSTTService(STTService):
|
|||||||
if not self._thread_task:
|
if not self._thread_task:
|
||||||
self._thread_task = self.create_task(self._thread_task_handler())
|
self._thread_task = self.create_task(self._thread_task_handler())
|
||||||
|
|
||||||
if not self._response_task:
|
|
||||||
self._response_queue = asyncio.Queue()
|
|
||||||
self._response_task = self.create_task(self._response_task_handler())
|
|
||||||
|
|
||||||
logger.debug(f"Initialized NvidiaSTTService with model: {self.model_name}")
|
logger.debug(f"Initialized NvidiaSTTService with model: {self.model_name}")
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
@@ -280,10 +275,6 @@ class NvidiaSTTService(STTService):
|
|||||||
await self.cancel_task(self._thread_task)
|
await self.cancel_task(self._thread_task)
|
||||||
self._thread_task = None
|
self._thread_task = None
|
||||||
|
|
||||||
if self._response_task:
|
|
||||||
await self.cancel_task(self._response_task)
|
|
||||||
self._response_task = None
|
|
||||||
|
|
||||||
def _response_handler(self):
|
def _response_handler(self):
|
||||||
responses = self._asr_service.streaming_response_generator(
|
responses = self._asr_service.streaming_response_generator(
|
||||||
audio_chunks=self,
|
audio_chunks=self,
|
||||||
@@ -292,9 +283,7 @@ class NvidiaSTTService(STTService):
|
|||||||
for response in responses:
|
for response in responses:
|
||||||
if not response.results:
|
if not response.results:
|
||||||
continue
|
continue
|
||||||
asyncio.run_coroutine_threadsafe(
|
asyncio.run_coroutine_threadsafe(self._handle_response(response), self.get_event_loop())
|
||||||
self._response_queue.put(response), self.get_event_loop()
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _thread_task_handler(self):
|
async def _thread_task_handler(self):
|
||||||
try:
|
try:
|
||||||
@@ -346,12 +335,6 @@ class NvidiaSTTService(STTService):
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
async def _response_task_handler(self):
|
|
||||||
while True:
|
|
||||||
response = await self._response_queue.get()
|
|
||||||
await self._handle_response(response)
|
|
||||||
self._response_queue.task_done()
|
|
||||||
|
|
||||||
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
||||||
"""Process audio data for speech-to-text transcription.
|
"""Process audio data for speech-to-text transcription.
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ gRPC API for high-quality speech synthesis.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import os
|
import os
|
||||||
from typing import AsyncGenerator, Mapping, Optional
|
from typing import AsyncGenerator, AsyncIterable, Generator, Mapping, Optional
|
||||||
|
|
||||||
from pipecat.utils.tracing.service_decorators import traced_tts
|
from pipecat.utils.tracing.service_decorators import traced_tts
|
||||||
|
|
||||||
@@ -35,14 +35,12 @@ from pipecat.transcriptions.language import Language
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
import riva.client
|
import riva.client
|
||||||
|
import riva.client.proto.riva_tts_pb2 as rtts
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
logger.error(f"Exception: {e}")
|
logger.error(f"Exception: {e}")
|
||||||
logger.error("In order to use NVIDIA Riva TTS, you need to `pip install pipecat-ai[nvidia]`.")
|
logger.error("In order to use NVIDIA Riva TTS, you need to `pip install pipecat-ai[nvidia]`.")
|
||||||
raise Exception(f"Missing module: {e}")
|
raise Exception(f"Missing module: {e}")
|
||||||
|
|
||||||
NVIDIA_TTS_TIMEOUT_SECS = 5
|
|
||||||
|
|
||||||
|
|
||||||
class NvidiaTTSService(TTSService):
|
class NvidiaTTSService(TTSService):
|
||||||
"""NVIDIA Riva text-to-speech service.
|
"""NVIDIA Riva text-to-speech service.
|
||||||
@@ -165,26 +163,30 @@ class NvidiaTTSService(TTSService):
|
|||||||
Frame: Audio frames containing the synthesized speech data.
|
Frame: Audio frames containing the synthesized speech data.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def read_audio_responses(queue: asyncio.Queue):
|
def read_audio_responses() -> Generator[rtts.SynthesizeSpeechResponse, None, None]:
|
||||||
def add_response(r):
|
responses = self._service.synthesize_online(
|
||||||
asyncio.run_coroutine_threadsafe(queue.put(r), self.get_event_loop())
|
text,
|
||||||
|
self._voice_id,
|
||||||
|
self._language_code,
|
||||||
|
sample_rate_hz=self.sample_rate,
|
||||||
|
zero_shot_audio_prompt_file=None,
|
||||||
|
zero_shot_quality=self._quality,
|
||||||
|
custom_dictionary={},
|
||||||
|
)
|
||||||
|
return responses
|
||||||
|
|
||||||
|
def async_next(it):
|
||||||
try:
|
try:
|
||||||
responses = self._service.synthesize_online(
|
return next(it)
|
||||||
text,
|
except StopIteration:
|
||||||
self._voice_id,
|
return None
|
||||||
self._language_code,
|
|
||||||
sample_rate_hz=self.sample_rate,
|
async def async_iterator(iterator) -> AsyncIterable[rtts.SynthesizeSpeechResponse]:
|
||||||
zero_shot_audio_prompt_file=None,
|
while True:
|
||||||
zero_shot_quality=self._quality,
|
item = await asyncio.to_thread(async_next, iterator)
|
||||||
custom_dictionary={},
|
if item is None:
|
||||||
)
|
return
|
||||||
for r in responses:
|
yield item
|
||||||
add_response(r)
|
|
||||||
add_response(None)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{self} exception: {e}")
|
|
||||||
add_response(None)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
assert self._service is not None, "TTS service not initialized"
|
assert self._service is not None, "TTS service not initialized"
|
||||||
@@ -195,12 +197,9 @@ class NvidiaTTSService(TTSService):
|
|||||||
|
|
||||||
logger.debug(f"{self}: Generating TTS [{text}]")
|
logger.debug(f"{self}: Generating TTS [{text}]")
|
||||||
|
|
||||||
queue = asyncio.Queue()
|
responses = await asyncio.to_thread(read_audio_responses)
|
||||||
await asyncio.to_thread(read_audio_responses, queue)
|
|
||||||
|
|
||||||
# Wait for the thread to start.
|
async for resp in async_iterator(responses):
|
||||||
resp = await asyncio.wait_for(queue.get(), timeout=NVIDIA_TTS_TIMEOUT_SECS)
|
|
||||||
while resp:
|
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
frame = TTSAudioRawFrame(
|
frame = TTSAudioRawFrame(
|
||||||
audio=resp.audio,
|
audio=resp.audio,
|
||||||
@@ -208,7 +207,6 @@ class NvidiaTTSService(TTSService):
|
|||||||
num_channels=1,
|
num_channels=1,
|
||||||
)
|
)
|
||||||
yield frame
|
yield frame
|
||||||
resp = await asyncio.wait_for(queue.get(), timeout=NVIDIA_TTS_TIMEOUT_SECS)
|
|
||||||
|
|
||||||
await self.start_tts_usage_metrics(text)
|
await self.start_tts_usage_metrics(text)
|
||||||
yield TTSStoppedFrame()
|
yield TTSStoppedFrame()
|
||||||
|
|||||||
Reference in New Issue
Block a user