Addressing the comments left in the PR review.
This commit is contained in:
@@ -46,15 +46,15 @@ from pipecat.utils.tracing.service_decorators import traced_stt
|
|||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class NvidiaSageMakerWSSTTSettings(STTSettings):
|
class NvidiaSageMakerSTTSettings(STTSettings):
|
||||||
"""Settings for NvidiaSageMakerWebsocketSTTService.
|
"""Settings for NvidiaSageMakerSTTService.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
language: ISO-639-1 language code passed to NIM (e.g. ``en-US``).
|
language: ISO-639-1 language code passed to NIM (e.g. ``en-US``).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
class NvidiaSageMakerWebsocketSTTService(STTService):
|
class NvidiaSageMakerSTTService(STTService):
|
||||||
"""NVIDIA Nemotron ASR STT service using SageMaker bidirectional streaming.
|
"""NVIDIA Nemotron ASR STT service using SageMaker bidirectional streaming.
|
||||||
|
|
||||||
Maintains a persistent HTTP/2 bidi-stream session to the SageMaker endpoint
|
Maintains a persistent HTTP/2 bidi-stream session to the SageMaker endpoint
|
||||||
@@ -65,16 +65,16 @@ class NvidiaSageMakerWebsocketSTTService(STTService):
|
|||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
stt = NvidiaSageMakerWebsocketSTTService(
|
stt = NvidiaSageMakerSTTService(
|
||||||
endpoint_name=os.getenv("SAGEMAKER_ASR_ENDPOINT_NAME"),
|
endpoint_name=os.getenv("SAGEMAKER_ASR_ENDPOINT_NAME"),
|
||||||
region=os.getenv("AWS_REGION", "us-west-2"),
|
region=os.getenv("AWS_REGION", "us-west-2"),
|
||||||
settings=NvidiaSageMakerWebsocketSTTService.Settings(
|
settings=NvidiaSageMakerSTTService.Settings(
|
||||||
language="en-US",
|
language="en-US",
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
Settings = NvidiaSageMakerWSSTTSettings
|
Settings = NvidiaSageMakerSTTSettings
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -82,7 +82,7 @@ class NvidiaSageMakerWebsocketSTTService(STTService):
|
|||||||
endpoint_name: str,
|
endpoint_name: str,
|
||||||
region: str = "us-west-2",
|
region: str = "us-west-2",
|
||||||
sample_rate: int | None = None,
|
sample_rate: int | None = None,
|
||||||
settings: NvidiaSageMakerWSSTTSettings | None = None,
|
settings: NvidiaSageMakerSTTSettings | None = None,
|
||||||
ttfs_p99_latency: float | None = 1.5,
|
ttfs_p99_latency: float | None = 1.5,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
@@ -301,6 +301,8 @@ class NvidiaSageMakerWebsocketSTTService(STTService):
|
|||||||
delta,
|
delta,
|
||||||
self._user_id,
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
|
language=self._settings.language,
|
||||||
|
result=msg,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -313,7 +315,9 @@ class NvidiaSageMakerWebsocketSTTService(STTService):
|
|||||||
transcript,
|
transcript,
|
||||||
self._user_id,
|
self._user_id,
|
||||||
time_now_iso8601(),
|
time_now_iso8601(),
|
||||||
|
language=self._settings.language,
|
||||||
result=msg,
|
result=msg,
|
||||||
|
finalized=True,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
await self._handle_transcription(transcript, True)
|
await self._handle_transcription(transcript, True)
|
||||||
|
|||||||
@@ -35,7 +35,7 @@ from pipecat.utils.tracing.service_decorators import traced_tts
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class NvidiaSageMakerTTSSettings(TTSSettings):
|
class NvidiaSageMakerTTSSettings(TTSSettings):
|
||||||
"""Settings for NvidiaSageMakerHTTPTTSService.
|
"""Settings for NVIDIA SageMaker TTS services.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
voice: NIM voice name (e.g. ``Magpie-Multilingual.EN-US.Aria``).
|
voice: NIM voice name (e.g. ``Magpie-Multilingual.EN-US.Aria``).
|
||||||
@@ -79,7 +79,6 @@ class NvidiaSageMakerHTTPTTSService(TTSService):
|
|||||||
endpoint_name: Name of the deployed SageMaker endpoint.
|
endpoint_name: Name of the deployed SageMaker endpoint.
|
||||||
region: AWS region where the endpoint lives.
|
region: AWS region where the endpoint lives.
|
||||||
sample_rate: Output sample rate in Hz. Defaults to bot's pipeline rate.
|
sample_rate: Output sample rate in Hz. Defaults to bot's pipeline rate.
|
||||||
params: Deprecated — use ``settings`` instead.
|
|
||||||
settings: Runtime-updatable settings (voice, language).
|
settings: Runtime-updatable settings (voice, language).
|
||||||
**kwargs: Forwarded to :class:`TTSService`.
|
**kwargs: Forwarded to :class:`TTSService`.
|
||||||
"""
|
"""
|
||||||
@@ -218,17 +217,7 @@ class NvidiaSageMakerHTTPTTSService(TTSService):
|
|||||||
await self.start_tts_usage_metrics(text)
|
await self.start_tts_usage_metrics(text)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
class NvidiaSageMakerTTSService(InterruptibleTTSService):
|
||||||
class NvidiaSageMakerWSTTSSettings(TTSSettings):
|
|
||||||
"""Settings for NvidiaSageMakerWebsocketTTSService.
|
|
||||||
|
|
||||||
Parameters:
|
|
||||||
voice: NIM voice name (e.g. ``Magpie-Multilingual.EN-US.Aria``).
|
|
||||||
language: BCP-47 language code passed to NIM (e.g. ``en-US``).
|
|
||||||
"""
|
|
||||||
|
|
||||||
|
|
||||||
class NvidiaSageMakerWebsocketTTSService(InterruptibleTTSService):
|
|
||||||
"""NVIDIA Magpie TTS service using SageMaker bidirectional streaming.
|
"""NVIDIA Magpie TTS service using SageMaker bidirectional streaming.
|
||||||
|
|
||||||
Maintains a persistent HTTP/2 bidi-stream session to the SageMaker endpoint
|
Maintains a persistent HTTP/2 bidi-stream session to the SageMaker endpoint
|
||||||
@@ -238,17 +227,17 @@ class NvidiaSageMakerWebsocketTTSService(InterruptibleTTSService):
|
|||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
tts = NvidiaSageMakerWebsocketTTSService(
|
tts = NvidiaSageMakerTTSService(
|
||||||
endpoint_name=os.getenv("SAGEMAKER_MAGPIE_ENDPOINT_NAME"),
|
endpoint_name=os.getenv("SAGEMAKER_MAGPIE_ENDPOINT_NAME"),
|
||||||
region=os.getenv("AWS_REGION", "us-west-2"),
|
region=os.getenv("AWS_REGION", "us-west-2"),
|
||||||
settings=NvidiaSageMakerWebsocketTTSService.Settings(
|
settings=NvidiaSageMakerTTSService.Settings(
|
||||||
voice="Magpie-Multilingual.EN-US.Aria",
|
voice="Magpie-Multilingual.EN-US.Aria",
|
||||||
language="en-US",
|
language="en-US",
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
Settings = NvidiaSageMakerWSTTSSettings
|
Settings = NvidiaSageMakerTTSSettings
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -256,7 +245,7 @@ class NvidiaSageMakerWebsocketTTSService(InterruptibleTTSService):
|
|||||||
endpoint_name: str,
|
endpoint_name: str,
|
||||||
region: str = "us-west-2",
|
region: str = "us-west-2",
|
||||||
sample_rate: int | None = None,
|
sample_rate: int | None = None,
|
||||||
settings: NvidiaSageMakerWSTTSSettings | None = None,
|
settings: NvidiaSageMakerTTSSettings | None = None,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
"""Initialize the SageMaker WebSocket TTS service.
|
"""Initialize the SageMaker WebSocket TTS service.
|
||||||
@@ -507,4 +496,4 @@ class NvidiaSageMakerWebsocketTTSService(InterruptibleTTSService):
|
|||||||
yield None
|
yield None
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self}: TTS error: {e}")
|
logger.error(f"{self}: TTS error: {e}")
|
||||||
yield ErrorFrame(error=f"NvidiaSageMakerWebsocketTTSService error: {e}")
|
yield ErrorFrame(error=f"NvidiaSageMakerTTSService error: {e}")
|
||||||
|
|||||||
Reference in New Issue
Block a user