Update STT service settings
This commit is contained in:
@@ -22,9 +22,12 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
)
|
)
|
||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.aws.llm import AWSBedrockLLMService
|
from pipecat.services.aws.llm import AWSBedrockLLMService, AWSBedrockLLMSettings
|
||||||
from pipecat.services.deepgram.sagemaker.stt import DeepgramSageMakerSTTService
|
from pipecat.services.deepgram.sagemaker.stt import DeepgramSageMakerSTTService
|
||||||
from pipecat.services.deepgram.sagemaker.tts import DeepgramSageMakerTTSService
|
from pipecat.services.deepgram.sagemaker.tts import (
|
||||||
|
DeepgramSageMakerTTSService,
|
||||||
|
DeepgramSageMakerTTSSettings,
|
||||||
|
)
|
||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
from pipecat.transports.daily.transport import DailyParams
|
from pipecat.transports.daily.transport import DailyParams
|
||||||
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
||||||
@@ -69,14 +72,18 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
tts = DeepgramSageMakerTTSService(
|
tts = DeepgramSageMakerTTSService(
|
||||||
endpoint_name=os.getenv("SAGEMAKER_TTS_ENDPOINT_NAME"),
|
endpoint_name=os.getenv("SAGEMAKER_TTS_ENDPOINT_NAME"),
|
||||||
region=os.getenv("AWS_REGION"),
|
region=os.getenv("AWS_REGION"),
|
||||||
voice="aura-2-andromeda-en",
|
settings=DeepgramSageMakerTTSSettings(
|
||||||
|
voice="aura-2-andromeda-en",
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
llm = AWSBedrockLLMService(
|
llm = AWSBedrockLLMService(
|
||||||
aws_region=os.getenv("AWS_REGION"),
|
aws_region=os.getenv("AWS_REGION"),
|
||||||
model="us.amazon.nova-pro-v1:0",
|
|
||||||
params=AWSBedrockLLMService.InputParams(temperature=0.8),
|
|
||||||
system_instruction="You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be spoken aloud, so avoid special characters that can't easily be spoken, such as emojis or bullet points. Respond to what the user said in a creative and helpful way.",
|
system_instruction="You are a helpful LLM in a WebRTC call. Your goal is to demonstrate your capabilities in a succinct way. Your output will be spoken aloud, so avoid special characters that can't easily be spoken, such as emojis or bullet points. Respond to what the user said in a creative and helpful way.",
|
||||||
|
settings=AWSBedrockLLMSettings(
|
||||||
|
model="us.amazon.nova-pro-v1:0",
|
||||||
|
temperature=0.8,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
context = LLMContext()
|
context = LLMContext()
|
||||||
|
|||||||
@@ -23,8 +23,8 @@ from pipecat.processors.aggregators.llm_response_universal import (
|
|||||||
from pipecat.runner.types import RunnerArguments
|
from pipecat.runner.types import RunnerArguments
|
||||||
from pipecat.runner.utils import create_transport
|
from pipecat.runner.utils import create_transport
|
||||||
from pipecat.services.openai.llm import OpenAILLMService
|
from pipecat.services.openai.llm import OpenAILLMService
|
||||||
from pipecat.services.openai.stt import OpenAISTTService
|
from pipecat.services.openai.stt import OpenAISTTService, OpenAISTTSettings
|
||||||
from pipecat.services.openai.tts import OpenAITTSService
|
from pipecat.services.openai.tts import OpenAITTSService, OpenAITTSSettings
|
||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
from pipecat.transports.daily.transport import DailyParams
|
from pipecat.transports.daily.transport import DailyParams
|
||||||
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
||||||
@@ -54,11 +54,18 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
|
|
||||||
stt = OpenAISTTService(
|
stt = OpenAISTTService(
|
||||||
api_key=os.getenv("OPENAI_API_KEY"),
|
api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
model="gpt-4o-transcribe",
|
settings=OpenAISTTSettings(
|
||||||
prompt="Expect words related to dogs, such as breed names.",
|
model="gpt-4o-transcribe",
|
||||||
|
prompt="Expect words related to dogs, such as breed names.",
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
tts = OpenAITTSService(api_key=os.getenv("OPENAI_API_KEY"), voice="ballad")
|
tts = OpenAITTSService(
|
||||||
|
api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
|
settings=OpenAITTSSettings(
|
||||||
|
voice="ballad",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(
|
llm = OpenAILLMService(
|
||||||
api_key=os.getenv("OPENAI_API_KEY"),
|
api_key=os.getenv("OPENAI_API_KEY"),
|
||||||
|
|||||||
@@ -124,6 +124,9 @@ AnyMessage = BeginMessage | TurnMessage | SpeechStartedMessage | TerminationMess
|
|||||||
class AssemblyAIConnectionParams(BaseModel):
|
class AssemblyAIConnectionParams(BaseModel):
|
||||||
"""Configuration parameters for AssemblyAI WebSocket connection.
|
"""Configuration parameters for AssemblyAI WebSocket connection.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=AssemblyAISTTSettings(foo=...)`` instead.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
sample_rate: Audio sample rate in Hz. Defaults to 16000.
|
sample_rate: Audio sample rate in Hz. Defaults to 16000.
|
||||||
encoding: Audio encoding format. Defaults to "pcm_s16le".
|
encoding: Audio encoding format. Defaults to "pcm_s16le".
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ WebSocket API for streaming audio transcription.
|
|||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, AsyncGenerator, Dict, Optional
|
from typing import Any, AsyncGenerator, Dict, List, Optional
|
||||||
from urllib.parse import urlencode
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -83,15 +83,38 @@ def map_language_from_assemblyai(language_code: str) -> Language:
|
|||||||
class AssemblyAISTTSettings(STTSettings):
|
class AssemblyAISTTSettings(STTSettings):
|
||||||
"""Settings for the AssemblyAI STT service.
|
"""Settings for the AssemblyAI STT service.
|
||||||
|
|
||||||
See :class:`AssemblyAIConnectionParams` for detailed parameter descriptions.
|
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
connection_params: Connection configuration parameters.
|
formatted_finals: Whether to enable transcript formatting.
|
||||||
|
word_finalization_max_wait_time: Maximum time to wait for word
|
||||||
|
finalization in milliseconds.
|
||||||
|
end_of_turn_confidence_threshold: Confidence threshold for
|
||||||
|
end-of-turn detection.
|
||||||
|
min_turn_silence: Minimum silence duration when confident about
|
||||||
|
end-of-turn.
|
||||||
|
max_turn_silence: Maximum silence duration before forcing
|
||||||
|
end-of-turn.
|
||||||
|
keyterms_prompt: List of key terms to guide transcription.
|
||||||
|
prompt: Optional text prompt to guide the transcription. Only
|
||||||
|
used when model is "u3-rt-pro".
|
||||||
|
language_detection: Enable automatic language detection.
|
||||||
|
format_turns: Whether to format transcript turns.
|
||||||
|
speaker_labels: Enable speaker diarization.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
connection_params: AssemblyAIConnectionParams | _NotGiven = field(
|
formatted_finals: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
word_finalization_max_wait_time: int | None | _NotGiven = field(
|
||||||
default_factory=lambda: NOT_GIVEN
|
default_factory=lambda: NOT_GIVEN
|
||||||
)
|
)
|
||||||
|
end_of_turn_confidence_threshold: float | None | _NotGiven = field(
|
||||||
|
default_factory=lambda: NOT_GIVEN
|
||||||
|
)
|
||||||
|
min_turn_silence: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
max_turn_silence: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
keyterms_prompt: List[str] | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
prompt: str | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
language_detection: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
format_turns: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
speaker_labels: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
|
||||||
|
|
||||||
class AssemblyAISTTService(WebsocketSTTService):
|
class AssemblyAISTTService(WebsocketSTTService):
|
||||||
@@ -110,6 +133,8 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
api_key: str,
|
api_key: str,
|
||||||
language: Optional[Language] = None,
|
language: Optional[Language] = None,
|
||||||
api_endpoint_base_url: str = "wss://streaming.assemblyai.com/v3/ws",
|
api_endpoint_base_url: str = "wss://streaming.assemblyai.com/v3/ws",
|
||||||
|
sample_rate: int = 16000,
|
||||||
|
encoding: str = "pcm_s16le",
|
||||||
connection_params: Optional[AssemblyAIConnectionParams] = None,
|
connection_params: Optional[AssemblyAIConnectionParams] = None,
|
||||||
vad_force_turn_endpoint: bool = True,
|
vad_force_turn_endpoint: bool = True,
|
||||||
should_interrupt: bool = True,
|
should_interrupt: bool = True,
|
||||||
@@ -123,8 +148,18 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
Args:
|
Args:
|
||||||
api_key: AssemblyAI API key for authentication.
|
api_key: AssemblyAI API key for authentication.
|
||||||
language: Language code for transcription. Defaults to English (Language.EN).
|
language: Language code for transcription. Defaults to English (Language.EN).
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=AssemblyAISTTSettings(language=...)`` instead.
|
||||||
|
|
||||||
api_endpoint_base_url: WebSocket endpoint URL. Defaults to AssemblyAI's streaming endpoint.
|
api_endpoint_base_url: WebSocket endpoint URL. Defaults to AssemblyAI's streaming endpoint.
|
||||||
connection_params: Connection configuration parameters. Defaults to AssemblyAIConnectionParams().
|
sample_rate: Audio sample rate in Hz. Defaults to 16000.
|
||||||
|
encoding: Audio encoding format. Defaults to "pcm_s16le".
|
||||||
|
connection_params: Connection configuration parameters.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=AssemblyAISTTSettings(...)`` instead.
|
||||||
|
|
||||||
vad_force_turn_endpoint: Controls turn detection mode.
|
vad_force_turn_endpoint: Controls turn detection mode.
|
||||||
When True (Pipecat mode, default): Forces AssemblyAI to return finals ASAP
|
When True (Pipecat mode, default): Forces AssemblyAI to return finals ASAP
|
||||||
so Pipecat's turn detection (e.g., Smart Turn) decides when the user is done.
|
so Pipecat's turn detection (e.g., Smart Turn) decides when the user is done.
|
||||||
@@ -135,7 +170,6 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
When False (AssemblyAI turn detection mode, u3-rt-pro only): AssemblyAI's model
|
When False (AssemblyAI turn detection mode, u3-rt-pro only): AssemblyAI's model
|
||||||
controls turn endings using built-in turn detection.
|
controls turn endings using built-in turn detection.
|
||||||
- Uses AssemblyAI API defaults for all parameters (unless user explicitly sets them)
|
- Uses AssemblyAI API defaults for all parameters (unless user explicitly sets them)
|
||||||
- Respects all user-provided connection_params as-is
|
|
||||||
- Emits UserStarted/StoppedSpeakingFrame from STT
|
- Emits UserStarted/StoppedSpeakingFrame from STT
|
||||||
- No ForceEndpoint on VAD stop
|
- No ForceEndpoint on VAD stop
|
||||||
should_interrupt: Whether to interrupt the bot when the user starts speaking
|
should_interrupt: Whether to interrupt the bot when the user starts speaking
|
||||||
@@ -145,39 +179,80 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
Use {speaker} for speaker label and {text} for transcript text.
|
Use {speaker} for speaker label and {text} for transcript text.
|
||||||
Example: "<{speaker}>{text}</{speaker}>" or "{speaker}: {text}"
|
Example: "<{speaker}>{text}</{speaker}>" or "{speaker}: {text}"
|
||||||
If None, transcript text is not modified. Defaults to None.
|
If None, transcript text is not modified. Defaults to None.
|
||||||
settings: Runtime-updatable settings. When provided alongside other
|
settings: Runtime-updatable settings. When provided alongside deprecated
|
||||||
parameters, ``settings`` values take precedence.
|
parameters, ``settings`` values take precedence.
|
||||||
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
||||||
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
||||||
**kwargs: Additional arguments passed to parent STTService class.
|
**kwargs: Additional arguments passed to parent STTService class.
|
||||||
"""
|
"""
|
||||||
# Resolve connection_params early — needed for validation and turn mode config
|
# 1. Initialize default_settings with hardcoded defaults
|
||||||
_connection_params = connection_params or AssemblyAIConnectionParams()
|
default_settings = AssemblyAISTTSettings(
|
||||||
|
model="u3-rt-pro",
|
||||||
|
language=Language.EN,
|
||||||
|
formatted_finals=True,
|
||||||
|
word_finalization_max_wait_time=None,
|
||||||
|
end_of_turn_confidence_threshold=None,
|
||||||
|
min_turn_silence=None,
|
||||||
|
max_turn_silence=None,
|
||||||
|
keyterms_prompt=None,
|
||||||
|
prompt=None,
|
||||||
|
language_detection=None,
|
||||||
|
format_turns=True,
|
||||||
|
speaker_labels=None,
|
||||||
|
)
|
||||||
|
|
||||||
# AssemblyAI turn detection mode (vad_force_turn_endpoint=False) requires the
|
# 2. Apply direct init arg overrides (deprecated)
|
||||||
# SpeechStarted event for reliable barge-in. Only u3-rt-pro supports
|
if language is not None:
|
||||||
# this. Other models must use Pipecat turn detection.
|
_warn_deprecated_param("language", AssemblyAISTTSettings, "language")
|
||||||
is_u3_pro = _connection_params.speech_model == "u3-rt-pro"
|
default_settings.language = language
|
||||||
|
|
||||||
|
# 3. Apply connection_params overrides (deprecated) — only if settings not provided
|
||||||
|
if connection_params is not None:
|
||||||
|
_warn_deprecated_param("connection_params", AssemblyAISTTSettings)
|
||||||
|
if not settings:
|
||||||
|
sample_rate = connection_params.sample_rate
|
||||||
|
encoding = connection_params.encoding
|
||||||
|
default_settings.model = connection_params.speech_model
|
||||||
|
default_settings.formatted_finals = connection_params.formatted_finals
|
||||||
|
default_settings.word_finalization_max_wait_time = (
|
||||||
|
connection_params.word_finalization_max_wait_time
|
||||||
|
)
|
||||||
|
default_settings.end_of_turn_confidence_threshold = (
|
||||||
|
connection_params.end_of_turn_confidence_threshold
|
||||||
|
)
|
||||||
|
default_settings.min_turn_silence = connection_params.min_turn_silence
|
||||||
|
default_settings.max_turn_silence = connection_params.max_turn_silence
|
||||||
|
default_settings.keyterms_prompt = connection_params.keyterms_prompt
|
||||||
|
default_settings.prompt = connection_params.prompt
|
||||||
|
default_settings.language_detection = connection_params.language_detection
|
||||||
|
default_settings.format_turns = connection_params.format_turns
|
||||||
|
default_settings.speaker_labels = connection_params.speaker_labels
|
||||||
|
|
||||||
|
# 4. Apply settings delta (canonical API, always wins)
|
||||||
|
if settings is not None:
|
||||||
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
|
# 5. Validate final settings
|
||||||
|
is_u3_pro = default_settings.model == "u3-rt-pro"
|
||||||
if not vad_force_turn_endpoint and not is_u3_pro:
|
if not vad_force_turn_endpoint and not is_u3_pro:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"AssemblyAI turn detection mode (vad_force_turn_endpoint=False) requires "
|
f"AssemblyAI turn detection mode (vad_force_turn_endpoint=False) requires "
|
||||||
f"u3-rt-pro for SpeechStarted support. Either set "
|
f"u3-rt-pro for SpeechStarted support. Either set "
|
||||||
f"vad_force_turn_endpoint=True for {_connection_params.speech_model}, "
|
f"vad_force_turn_endpoint=True for {default_settings.model}, "
|
||||||
f"or use speech_model='u3-rt-pro'."
|
f"or use model='u3-rt-pro'."
|
||||||
)
|
)
|
||||||
|
|
||||||
# Validate that prompt and keyterms_prompt are not both set
|
if default_settings.prompt is not None and default_settings.keyterms_prompt is not None:
|
||||||
if _connection_params.prompt is not None and _connection_params.keyterms_prompt is not None:
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"The prompt and keyterms_prompt parameters cannot be used in the same request. "
|
"The prompt and keyterms_prompt parameters cannot be used in the same request. "
|
||||||
"Please choose either one or the other based on your use case. When you use "
|
"Please choose either one or the other based on your use case. When you use "
|
||||||
"keyterms_prompt, your boosted words are appended to the default prompt automatically. "
|
"keyterms_prompt, your boosted words are appended to the default prompt automatically. "
|
||||||
"Or to boost within prompt: <prompt> + Make sure to boost the words <keyterms> in the audio. "
|
"Or to boost within prompt: <prompt> + Make sure to boost the words <keyterms> "
|
||||||
|
"in the audio. "
|
||||||
"For more info go to: https://www.assemblyai.com/docs/streaming/universal-3-pro"
|
"For more info go to: https://www.assemblyai.com/docs/streaming/universal-3-pro"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Warn if user sets a custom prompt (recommend testing without one first)
|
if default_settings.prompt is not None:
|
||||||
if _connection_params.prompt is not None:
|
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"Custom prompt detected. Prompting is a beta feature. We recommend testing "
|
"Custom prompt detected. Prompting is a beta feature. We recommend testing "
|
||||||
"with no prompt first, as this will use our optimized default prompt for "
|
"with no prompt first, as this will use our optimized default prompt for "
|
||||||
@@ -186,35 +261,12 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
"https://www.assemblyai.com/docs/streaming/prompting"
|
"https://www.assemblyai.com/docs/streaming/prompting"
|
||||||
)
|
)
|
||||||
|
|
||||||
# When vad_force_turn_endpoint is enabled, configure connection params
|
# 6. Configure pipecat turn mode (mutates default_settings)
|
||||||
# for Pipecat turn detection mode (fast finals for smart turn analyzer)
|
|
||||||
if vad_force_turn_endpoint:
|
if vad_force_turn_endpoint:
|
||||||
_connection_params = self._configure_pipecat_turn_mode(_connection_params, is_u3_pro)
|
self._configure_pipecat_turn_mode(default_settings, is_u3_pro)
|
||||||
|
|
||||||
# 1. Initialize default_settings with hardcoded defaults
|
|
||||||
default_settings = AssemblyAISTTSettings(
|
|
||||||
model=None,
|
|
||||||
language=Language.EN,
|
|
||||||
connection_params=AssemblyAIConnectionParams(),
|
|
||||||
)
|
|
||||||
|
|
||||||
# 2. Apply direct init arg overrides (deprecated)
|
|
||||||
if language is not None:
|
|
||||||
_warn_deprecated_param("language", AssemblyAISTTSettings, "language")
|
|
||||||
default_settings.language = language
|
|
||||||
|
|
||||||
# 3. Apply connection_params overrides — only if settings not provided
|
|
||||||
if connection_params is not None:
|
|
||||||
_warn_deprecated_param("connection_params", AssemblyAISTTSettings)
|
|
||||||
if not settings:
|
|
||||||
default_settings.connection_params = _connection_params
|
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
|
||||||
if settings is not None:
|
|
||||||
default_settings.apply_update(settings)
|
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
sample_rate=_connection_params.sample_rate,
|
sample_rate=sample_rate,
|
||||||
ttfs_p99_latency=ttfs_p99_latency,
|
ttfs_p99_latency=ttfs_p99_latency,
|
||||||
settings=default_settings,
|
settings=default_settings,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
@@ -226,6 +278,9 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
self._should_interrupt = should_interrupt
|
self._should_interrupt = should_interrupt
|
||||||
self._speaker_format = speaker_format
|
self._speaker_format = speaker_format
|
||||||
|
|
||||||
|
# Init-only audio config (not runtime-updatable)
|
||||||
|
self._encoding = encoding
|
||||||
|
|
||||||
self._termination_event = asyncio.Event()
|
self._termination_event = asyncio.Event()
|
||||||
self._received_termination = False
|
self._received_termination = False
|
||||||
self._connected = False
|
self._connected = False
|
||||||
@@ -238,10 +293,8 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
|
|
||||||
self._user_speaking = False
|
self._user_speaking = False
|
||||||
|
|
||||||
def _configure_pipecat_turn_mode(
|
def _configure_pipecat_turn_mode(self, settings: AssemblyAISTTSettings, is_u3_pro: bool):
|
||||||
self, connection_params: AssemblyAIConnectionParams, is_u3_pro: bool
|
"""Configure settings for Pipecat turn detection mode.
|
||||||
) -> AssemblyAIConnectionParams:
|
|
||||||
"""Configure connection params for Pipecat turn detection mode.
|
|
||||||
|
|
||||||
When vad_force_turn_endpoint is enabled, force AssemblyAI to return
|
When vad_force_turn_endpoint is enabled, force AssemblyAI to return
|
||||||
finals as fast as possible so Pipecat's smart turn analyzer can decide
|
finals as fast as possible so Pipecat's smart turn analyzer can decide
|
||||||
@@ -260,46 +313,31 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
- max_turn_silence: not set (API default)
|
- max_turn_silence: not set (API default)
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
connection_params: The user-provided connection parameters.
|
settings: The settings to configure in place.
|
||||||
is_u3_pro: Whether using u3-rt-pro model.
|
is_u3_pro: Whether using u3-rt-pro model.
|
||||||
|
|
||||||
Returns:
|
|
||||||
Updated connection parameters configured for Pipecat turn mode.
|
|
||||||
"""
|
"""
|
||||||
updates = {}
|
|
||||||
|
|
||||||
if is_u3_pro:
|
if is_u3_pro:
|
||||||
# u3-rt-pro: Synchronize max_turn_silence with min_turn_silence
|
# u3-rt-pro: Synchronize max_turn_silence with min_turn_silence
|
||||||
min_silence = connection_params.min_turn_silence
|
min_silence = settings.min_turn_silence
|
||||||
if min_silence is None:
|
if min_silence is None:
|
||||||
min_silence = 100
|
min_silence = 100
|
||||||
|
|
||||||
# Warn if user set max_turn_silence (will be overridden)
|
# Warn if user set max_turn_silence (will be overridden)
|
||||||
if connection_params.max_turn_silence is not None:
|
if settings.max_turn_silence is not None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Your max_turn_silence value ({connection_params.max_turn_silence}ms) will be "
|
f"Your max_turn_silence value ({settings.max_turn_silence}ms) will be "
|
||||||
f"OVERRIDDEN in Pipecat mode (vad_force_turn_endpoint=True). It will be set to "
|
f"OVERRIDDEN in Pipecat mode (vad_force_turn_endpoint=True). It will be set to "
|
||||||
f"{min_silence}ms (matching min_turn_silence) and SENT to "
|
f"{min_silence}ms (matching min_turn_silence) and SENT to "
|
||||||
f"AssemblyAI to avoid double turn detection. To use your max_turn_silence as-is, "
|
f"AssemblyAI to avoid double turn detection. To use your max_turn_silence as-is, "
|
||||||
f"switch to AssemblyAI turn detection mode (vad_force_turn_endpoint=False)."
|
f"switch to AssemblyAI turn detection mode (vad_force_turn_endpoint=False)."
|
||||||
)
|
)
|
||||||
|
|
||||||
updates = {
|
settings.min_turn_silence = min_silence
|
||||||
"min_turn_silence": min_silence,
|
settings.max_turn_silence = min_silence
|
||||||
"max_turn_silence": min_silence,
|
|
||||||
}
|
|
||||||
else:
|
else:
|
||||||
# universal-streaming: Different configuration (works differently)
|
# universal-streaming: Different configuration (works differently)
|
||||||
updates = {
|
settings.end_of_turn_confidence_threshold = 1.0
|
||||||
"end_of_turn_confidence_threshold": 1.0,
|
settings.min_turn_silence = 160
|
||||||
"min_turn_silence": 160,
|
|
||||||
}
|
|
||||||
|
|
||||||
# Apply updates if any
|
|
||||||
if updates:
|
|
||||||
connection_params = connection_params.model_copy(update=updates)
|
|
||||||
|
|
||||||
return connection_params
|
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
"""Check if the service can generate metrics.
|
"""Check if the service can generate metrics.
|
||||||
@@ -309,18 +347,11 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
"""
|
"""
|
||||||
return True
|
return True
|
||||||
|
|
||||||
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
async def _update_settings(self, delta: AssemblyAISTTSettings) -> dict[str, Any]:
|
||||||
"""Apply a settings delta and send UpdateConfiguration if connected.
|
"""Apply a settings delta and reconnect to apply changes.
|
||||||
|
|
||||||
Stores settings changes and sends UpdateConfiguration message to AssemblyAI
|
|
||||||
without reconnecting. Supports updating:
|
|
||||||
- keyterms_prompt: List of terms to boost (can be empty array to clear)
|
|
||||||
- prompt: Custom prompt text (u3-rt-pro only)
|
|
||||||
- max_turn_silence: Maximum silence before forcing turn end
|
|
||||||
- min_turn_silence: Silence before EOT check
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
delta: A :class:`STTSettings` (or ``AssemblyAISTTSettings``) delta.
|
delta: A settings delta with updated values.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dict mapping changed field names to their previous values.
|
Dict mapping changed field names to their previous values.
|
||||||
@@ -330,72 +361,9 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
if not changed:
|
if not changed:
|
||||||
return changed
|
return changed
|
||||||
|
|
||||||
# If websocket is connected, send UpdateConfiguration for supported params
|
# Reconnect to apply updated settings (they become WS query params)
|
||||||
if (
|
await self._disconnect()
|
||||||
self._websocket
|
await self._connect()
|
||||||
and self._websocket.state is State.OPEN
|
|
||||||
and "connection_params" in changed
|
|
||||||
):
|
|
||||||
# Build UpdateConfiguration message
|
|
||||||
update_config = {"type": "UpdateConfiguration"}
|
|
||||||
conn_params = self._settings.connection_params
|
|
||||||
|
|
||||||
# Get the old connection_params to see what changed
|
|
||||||
old_conn_params = changed.get("connection_params")
|
|
||||||
|
|
||||||
# Check each potentially changed parameter
|
|
||||||
if (
|
|
||||||
old_conn_params is None
|
|
||||||
or conn_params.keyterms_prompt != old_conn_params.keyterms_prompt
|
|
||||||
):
|
|
||||||
if conn_params.keyterms_prompt is not None:
|
|
||||||
update_config["keyterms_prompt"] = conn_params.keyterms_prompt
|
|
||||||
logger.info(f"Updating keyterms_prompt to: {conn_params.keyterms_prompt}")
|
|
||||||
|
|
||||||
if old_conn_params is None or conn_params.prompt != old_conn_params.prompt:
|
|
||||||
if conn_params.prompt is not None:
|
|
||||||
if conn_params.speech_model != "u3-rt-pro":
|
|
||||||
logger.warning(
|
|
||||||
f"prompt parameter is only supported with u3-rt-pro model, "
|
|
||||||
f"current model is {conn_params.speech_model}"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
update_config["prompt"] = conn_params.prompt
|
|
||||||
logger.info(f"Updating prompt")
|
|
||||||
|
|
||||||
if (
|
|
||||||
old_conn_params is None
|
|
||||||
or conn_params.max_turn_silence != old_conn_params.max_turn_silence
|
|
||||||
):
|
|
||||||
if conn_params.max_turn_silence is not None:
|
|
||||||
update_config["max_turn_silence"] = conn_params.max_turn_silence
|
|
||||||
logger.info(f"Updating max_turn_silence to: {conn_params.max_turn_silence}ms")
|
|
||||||
|
|
||||||
if (
|
|
||||||
old_conn_params is None
|
|
||||||
or conn_params.min_turn_silence != old_conn_params.min_turn_silence
|
|
||||||
):
|
|
||||||
if conn_params.min_turn_silence is not None:
|
|
||||||
update_config["min_turn_silence"] = conn_params.min_turn_silence
|
|
||||||
logger.info(f"Updating min_turn_silence to: {conn_params.min_turn_silence}ms")
|
|
||||||
|
|
||||||
# Send update if we have parameters to update
|
|
||||||
if len(update_config) > 1: # More than just "type"
|
|
||||||
try:
|
|
||||||
await self._websocket.send(json.dumps(update_config))
|
|
||||||
logger.info(f"Sent UpdateConfiguration: {update_config}")
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"Failed to send UpdateConfiguration: {e}")
|
|
||||||
elif "connection_params" in changed:
|
|
||||||
logger.warning(
|
|
||||||
"Connection params changed but WebSocket not connected. "
|
|
||||||
"Settings will be applied on next connection."
|
|
||||||
)
|
|
||||||
|
|
||||||
# Warn about other settings that can't be changed dynamically
|
|
||||||
other_changes = {k: v for k, v in changed.items() if k not in ["connection_params"]}
|
|
||||||
if other_changes:
|
|
||||||
self._warn_unhandled_updated_settings(other_changes)
|
|
||||||
|
|
||||||
return changed
|
return changed
|
||||||
|
|
||||||
@@ -473,19 +441,41 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
|
|
||||||
def _build_ws_url(self) -> str:
|
def _build_ws_url(self) -> str:
|
||||||
"""Build WebSocket URL with query parameters using urllib.parse.urlencode."""
|
"""Build WebSocket URL with query parameters using urllib.parse.urlencode."""
|
||||||
params = {}
|
s = self._settings
|
||||||
for k, v in self._settings.connection_params.model_dump().items():
|
params: dict[str, Any] = {}
|
||||||
# Skip deprecated parameter - it's been migrated to min_turn_silence
|
|
||||||
if k == "min_end_of_turn_silence_when_confident":
|
# Init-only audio config
|
||||||
continue
|
params["sample_rate"] = self.sample_rate
|
||||||
|
params["encoding"] = self._encoding
|
||||||
|
|
||||||
|
# Map model → speech_model (AssemblyAI API naming)
|
||||||
|
if s.model is not None:
|
||||||
|
params["speech_model"] = s.model
|
||||||
|
|
||||||
|
# Settings fields (skip None values)
|
||||||
|
optional_fields = {
|
||||||
|
"formatted_finals": s.formatted_finals,
|
||||||
|
"word_finalization_max_wait_time": s.word_finalization_max_wait_time,
|
||||||
|
"end_of_turn_confidence_threshold": s.end_of_turn_confidence_threshold,
|
||||||
|
"min_turn_silence": s.min_turn_silence,
|
||||||
|
"max_turn_silence": s.max_turn_silence,
|
||||||
|
"prompt": s.prompt,
|
||||||
|
"language_detection": s.language_detection,
|
||||||
|
"format_turns": s.format_turns,
|
||||||
|
"speaker_labels": s.speaker_labels,
|
||||||
|
}
|
||||||
|
|
||||||
|
for k, v in optional_fields.items():
|
||||||
if v is not None:
|
if v is not None:
|
||||||
if k == "keyterms_prompt":
|
if isinstance(v, bool):
|
||||||
params[k] = json.dumps(v)
|
|
||||||
elif isinstance(v, bool):
|
|
||||||
params[k] = str(v).lower()
|
params[k] = str(v).lower()
|
||||||
else:
|
else:
|
||||||
params[k] = v
|
params[k] = v
|
||||||
|
|
||||||
|
# Special handling for keyterms_prompt (needs JSON encoding)
|
||||||
|
if s.keyterms_prompt is not None:
|
||||||
|
params["keyterms_prompt"] = json.dumps(s.keyterms_prompt)
|
||||||
|
|
||||||
if params:
|
if params:
|
||||||
query_string = urlencode(params)
|
query_string = urlencode(params)
|
||||||
return f"{self._api_endpoint_base_url}?{query_string}"
|
return f"{self._api_endpoint_base_url}?{query_string}"
|
||||||
@@ -717,7 +707,7 @@ class AssemblyAISTTService(WebsocketSTTService):
|
|||||||
|
|
||||||
# Determine if this is a final turn from AssemblyAI
|
# Determine if this is a final turn from AssemblyAI
|
||||||
is_final_turn = message.end_of_turn and (
|
is_final_turn = message.end_of_turn and (
|
||||||
not self._settings.connection_params.format_turns or message.turn_is_formatted
|
not self._settings.format_turns or message.turn_is_formatted
|
||||||
)
|
)
|
||||||
|
|
||||||
if self._vad_force_turn_endpoint:
|
if self._vad_force_turn_endpoint:
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ import json
|
|||||||
import os
|
import os
|
||||||
import random
|
import random
|
||||||
import string
|
import string
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, Optional
|
from typing import Any, AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -29,7 +29,7 @@ from pipecat.frames.frames import (
|
|||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
)
|
)
|
||||||
from pipecat.services.aws.utils import build_event_message, decode_event, get_presigned_url
|
from pipecat.services.aws.utils import build_event_message, decode_event, get_presigned_url
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.settings import STTSettings, _warn_deprecated_param
|
||||||
from pipecat.services.stt_latency import AWS_TRANSCRIBE_TTFS_P99
|
from pipecat.services.stt_latency import AWS_TRANSCRIBE_TTFS_P99
|
||||||
from pipecat.services.stt_service import WebsocketSTTService
|
from pipecat.services.stt_service import WebsocketSTTService
|
||||||
from pipecat.transcriptions.language import Language, resolve_language
|
from pipecat.transcriptions.language import Language, resolve_language
|
||||||
@@ -47,21 +47,9 @@ except ModuleNotFoundError as e:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AWSTranscribeSTTSettings(STTSettings):
|
class AWSTranscribeSTTSettings(STTSettings):
|
||||||
"""Settings for the AWS Transcribe STT service.
|
"""Settings for the AWS Transcribe STT service."""
|
||||||
|
|
||||||
Parameters:
|
pass
|
||||||
sample_rate: Audio sample rate in Hz (8000 or 16000).
|
|
||||||
media_encoding: Audio encoding format (e.g. "linear16").
|
|
||||||
number_of_channels: Number of audio channels.
|
|
||||||
show_speaker_label: Whether to show speaker labels.
|
|
||||||
enable_channel_identification: Whether to enable channel identification.
|
|
||||||
"""
|
|
||||||
|
|
||||||
sample_rate: int | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
media_encoding: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
number_of_channels: int | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
show_speaker_label: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
enable_channel_identification: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
|
|
||||||
|
|
||||||
class AWSTranscribeSTTService(WebsocketSTTService):
|
class AWSTranscribeSTTService(WebsocketSTTService):
|
||||||
@@ -94,11 +82,9 @@ class AWSTranscribeSTTService(WebsocketSTTService):
|
|||||||
aws_access_key_id: AWS access key ID. If None, uses AWS_ACCESS_KEY_ID environment variable.
|
aws_access_key_id: AWS access key ID. If None, uses AWS_ACCESS_KEY_ID environment variable.
|
||||||
aws_session_token: AWS session token for temporary credentials. If None, uses AWS_SESSION_TOKEN environment variable.
|
aws_session_token: AWS session token for temporary credentials. If None, uses AWS_SESSION_TOKEN environment variable.
|
||||||
region: AWS region for the service.
|
region: AWS region for the service.
|
||||||
sample_rate: Audio sample rate in Hz. Must be 8000 or 16000.
|
sample_rate: Audio sample rate in Hz. If None, uses the pipeline sample rate.
|
||||||
|
AWS Transcribe only supports 8000 or 16000 Hz; other values are
|
||||||
.. deprecated:: 0.0.105
|
clamped to 16000 Hz at connect time.
|
||||||
Use ``settings=AWSTranscribeSTTSettings(sample_rate=...)`` instead.
|
|
||||||
|
|
||||||
language: Language for transcription.
|
language: Language for transcription.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
@@ -113,17 +99,9 @@ class AWSTranscribeSTTService(WebsocketSTTService):
|
|||||||
# 1. Initialize default_settings with hardcoded defaults
|
# 1. Initialize default_settings with hardcoded defaults
|
||||||
default_settings = AWSTranscribeSTTSettings(
|
default_settings = AWSTranscribeSTTSettings(
|
||||||
language=self.language_to_service_language(Language.EN) or "en-US",
|
language=self.language_to_service_language(Language.EN) or "en-US",
|
||||||
sample_rate=16000,
|
|
||||||
media_encoding="linear16",
|
|
||||||
number_of_channels=1,
|
|
||||||
show_speaker_label=False,
|
|
||||||
enable_channel_identification=False,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Apply direct init arg overrides (deprecated)
|
# 2. Apply direct init arg overrides (deprecated)
|
||||||
if sample_rate is not None:
|
|
||||||
_warn_deprecated_param("sample_rate", AWSTranscribeSTTSettings, "sample_rate")
|
|
||||||
default_settings.sample_rate = sample_rate
|
|
||||||
if language is not None:
|
if language is not None:
|
||||||
_warn_deprecated_param("language", AWSTranscribeSTTSettings, "language")
|
_warn_deprecated_param("language", AWSTranscribeSTTSettings, "language")
|
||||||
default_settings.language = self.language_to_service_language(language) or "en-US"
|
default_settings.language = self.language_to_service_language(language) or "en-US"
|
||||||
@@ -135,17 +113,17 @@ class AWSTranscribeSTTService(WebsocketSTTService):
|
|||||||
default_settings.apply_update(settings)
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
|
sample_rate=sample_rate,
|
||||||
ttfs_p99_latency=ttfs_p99_latency,
|
ttfs_p99_latency=ttfs_p99_latency,
|
||||||
settings=default_settings,
|
settings=default_settings,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Validate sample rate - AWS Transcribe only supports 8000 Hz or 16000 Hz
|
# Init-only connection config (not runtime-updatable).
|
||||||
if default_settings.sample_rate not in [8000, 16000]:
|
self._media_encoding = "linear16"
|
||||||
logger.warning(
|
self._number_of_channels = 1
|
||||||
f"AWS Transcribe only supports 8000 Hz or 16000 Hz sample rates. Converting from {default_settings.sample_rate} Hz to 16000 Hz."
|
self._show_speaker_label = False
|
||||||
)
|
self._enable_channel_identification = False
|
||||||
self._settings.sample_rate = 16000
|
|
||||||
|
|
||||||
self._credentials = {
|
self._credentials = {
|
||||||
"aws_access_key_id": aws_access_key_id or os.getenv("AWS_ACCESS_KEY_ID"),
|
"aws_access_key_id": aws_access_key_id or os.getenv("AWS_ACCESS_KEY_ID"),
|
||||||
@@ -293,6 +271,15 @@ class AWSTranscribeSTTService(WebsocketSTTService):
|
|||||||
if not language_code:
|
if not language_code:
|
||||||
raise ValueError(f"Unsupported language: {language_code}")
|
raise ValueError(f"Unsupported language: {language_code}")
|
||||||
|
|
||||||
|
# Validate sample rate — AWS Transcribe only supports 8000 or 16000 Hz
|
||||||
|
connect_sample_rate = self.sample_rate
|
||||||
|
if connect_sample_rate not in (8000, 16000):
|
||||||
|
logger.warning(
|
||||||
|
f"AWS Transcribe only supports 8000 Hz or 16000 Hz sample rates. "
|
||||||
|
f"Converting from {connect_sample_rate} Hz to 16000 Hz."
|
||||||
|
)
|
||||||
|
connect_sample_rate = 16000
|
||||||
|
|
||||||
# Generate random websocket key
|
# Generate random websocket key
|
||||||
websocket_key = "".join(
|
websocket_key = "".join(
|
||||||
random.choices(
|
random.choices(
|
||||||
@@ -318,14 +305,14 @@ class AWSTranscribeSTTService(WebsocketSTTService):
|
|||||||
},
|
},
|
||||||
language_code=language_code,
|
language_code=language_code,
|
||||||
media_encoding=self.get_service_encoding(
|
media_encoding=self.get_service_encoding(
|
||||||
self._settings.media_encoding
|
self._media_encoding
|
||||||
), # Convert to AWS format
|
), # Convert to AWS format
|
||||||
sample_rate=self._settings.sample_rate,
|
sample_rate=connect_sample_rate,
|
||||||
number_of_channels=self._settings.number_of_channels,
|
number_of_channels=self._number_of_channels,
|
||||||
enable_partial_results_stabilization=True,
|
enable_partial_results_stabilization=True,
|
||||||
partial_results_stability="high",
|
partial_results_stability="high",
|
||||||
show_speaker_label=self._settings.show_speaker_label,
|
show_speaker_label=self._show_speaker_label,
|
||||||
enable_channel_identification=self._settings.enable_channel_identification,
|
enable_channel_identification=self._enable_channel_identification,
|
||||||
)
|
)
|
||||||
|
|
||||||
logger.debug(f"{self} Connecting to WebSocket with URL: {presigned_url[:100]}...")
|
logger.debug(f"{self} Connecting to WebSocket with URL: {presigned_url[:100]}...")
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ Speech SDK for real-time audio transcription.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, Optional
|
from typing import Any, AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -26,7 +26,7 @@ from pipecat.frames.frames import (
|
|||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
)
|
)
|
||||||
from pipecat.services.azure.common import language_to_azure_language
|
from pipecat.services.azure.common import language_to_azure_language
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.settings import STTSettings, _warn_deprecated_param
|
||||||
from pipecat.services.stt_latency import AZURE_TTFS_P99
|
from pipecat.services.stt_latency import AZURE_TTFS_P99
|
||||||
from pipecat.services.stt_service import STTService
|
from pipecat.services.stt_service import STTService
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
@@ -53,15 +53,9 @@ except ModuleNotFoundError as e:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class AzureSTTSettings(STTSettings):
|
class AzureSTTSettings(STTSettings):
|
||||||
"""Settings for the Azure STT service.
|
"""Settings for the Azure STT service."""
|
||||||
|
|
||||||
Parameters:
|
pass
|
||||||
region: Azure region for the Speech service.
|
|
||||||
sample_rate: Audio sample rate in Hz.
|
|
||||||
"""
|
|
||||||
|
|
||||||
region: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
sample_rate: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
|
|
||||||
|
|
||||||
class AzureSTTService(STTService):
|
class AzureSTTService(STTService):
|
||||||
@@ -110,9 +104,7 @@ class AzureSTTService(STTService):
|
|||||||
# 1. Initialize default_settings with hardcoded defaults
|
# 1. Initialize default_settings with hardcoded defaults
|
||||||
default_settings = AzureSTTSettings(
|
default_settings = AzureSTTSettings(
|
||||||
model=None,
|
model=None,
|
||||||
region=region,
|
|
||||||
language=language_to_azure_language(Language.EN_US),
|
language=language_to_azure_language(Language.EN_US),
|
||||||
sample_rate=sample_rate,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Apply direct init arg overrides (deprecated)
|
# 2. Apply direct init arg overrides (deprecated)
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ the Cartesia Live transcription API for real-time speech recognition.
|
|||||||
|
|
||||||
import json
|
import json
|
||||||
import urllib.parse
|
import urllib.parse
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, Optional
|
from typing import Any, AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -28,7 +28,7 @@ from pipecat.frames.frames import (
|
|||||||
VADUserStoppedSpeakingFrame,
|
VADUserStoppedSpeakingFrame,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.settings import STTSettings, _warn_deprecated_param
|
||||||
from pipecat.services.stt_latency import CARTESIA_TTFS_P99
|
from pipecat.services.stt_latency import CARTESIA_TTFS_P99
|
||||||
from pipecat.services.stt_service import WebsocketSTTService
|
from pipecat.services.stt_service import WebsocketSTTService
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
@@ -46,20 +46,17 @@ except ModuleNotFoundError as e:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class CartesiaSTTSettings(STTSettings):
|
class CartesiaSTTSettings(STTSettings):
|
||||||
"""Settings for the Cartesia STT service.
|
"""Settings for the Cartesia STT service."""
|
||||||
|
|
||||||
Parameters:
|
pass
|
||||||
encoding: Audio encoding format (e.g. ``"pcm_s16le"``).
|
|
||||||
"""
|
|
||||||
|
|
||||||
encoding: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
|
|
||||||
|
|
||||||
class CartesiaLiveOptions:
|
class CartesiaLiveOptions:
|
||||||
"""Configuration options for Cartesia Live STT service.
|
"""Configuration options for Cartesia Live STT service.
|
||||||
|
|
||||||
Manages transcription parameters including model selection, language,
|
.. deprecated:: 0.0.105
|
||||||
audio encoding format, and sample rate settings.
|
Use ``settings=CartesiaSTTSettings(...)`` for model/language and
|
||||||
|
direct ``__init__`` parameters for encoding/sample_rate instead.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -156,7 +153,8 @@ class CartesiaSTTService(WebsocketSTTService):
|
|||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
base_url: str = "",
|
base_url: str = "",
|
||||||
sample_rate: int = 16000,
|
encoding: str = "pcm_s16le",
|
||||||
|
sample_rate: Optional[int] = None,
|
||||||
live_options: Optional[CartesiaLiveOptions] = None,
|
live_options: Optional[CartesiaLiveOptions] = None,
|
||||||
settings: Optional[CartesiaSTTSettings] = None,
|
settings: Optional[CartesiaSTTSettings] = None,
|
||||||
ttfs_p99_latency: Optional[float] = CARTESIA_TTFS_P99,
|
ttfs_p99_latency: Optional[float] = CARTESIA_TTFS_P99,
|
||||||
@@ -167,44 +165,42 @@ class CartesiaSTTService(WebsocketSTTService):
|
|||||||
Args:
|
Args:
|
||||||
api_key: Authentication key for Cartesia API.
|
api_key: Authentication key for Cartesia API.
|
||||||
base_url: Custom API endpoint URL. If empty, uses default.
|
base_url: Custom API endpoint URL. If empty, uses default.
|
||||||
sample_rate: Audio sample rate in Hz. Defaults to 16000.
|
encoding: Audio encoding format. Defaults to "pcm_s16le".
|
||||||
|
sample_rate: Audio sample rate in Hz. If None, uses the pipeline
|
||||||
|
sample rate.
|
||||||
live_options: Configuration options for transcription service.
|
live_options: Configuration options for transcription service.
|
||||||
settings: Runtime-updatable settings. When provided alongside
|
|
||||||
``live_options``, ``settings`` values take precedence.
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=CartesiaSTTSettings(...)`` for model/language
|
||||||
|
and direct init parameters for encoding/sample_rate instead.
|
||||||
|
|
||||||
|
settings: Runtime-updatable settings. When provided alongside deprecated
|
||||||
|
parameters, ``settings`` values take precedence.
|
||||||
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
||||||
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
||||||
**kwargs: Additional arguments passed to parent STTService.
|
**kwargs: Additional arguments passed to parent STTService.
|
||||||
"""
|
"""
|
||||||
sample_rate = sample_rate or (live_options.sample_rate if live_options else None)
|
|
||||||
|
|
||||||
# 1. Initialize default_settings with hardcoded defaults
|
# 1. Initialize default_settings with hardcoded defaults
|
||||||
default_settings = CartesiaSTTSettings(
|
default_settings = CartesiaSTTSettings(
|
||||||
model="ink-whisper",
|
model="ink-whisper",
|
||||||
language=Language.EN.value,
|
language=Language.EN.value,
|
||||||
encoding="pcm_s16le",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. (no deprecated direct args for this service)
|
# 2. Apply live_options overrides — only if settings not provided
|
||||||
|
|
||||||
# 3. Apply live_options overrides — only if settings not provided
|
|
||||||
if live_options is not None:
|
if live_options is not None:
|
||||||
_warn_deprecated_param("live_options", CartesiaSTTSettings)
|
_warn_deprecated_param("live_options", CartesiaSTTSettings)
|
||||||
if not settings:
|
if not settings:
|
||||||
lo_dict = live_options.to_dict()
|
if live_options.sample_rate and sample_rate is None:
|
||||||
# Filter out "None" string values
|
sample_rate = live_options.sample_rate
|
||||||
lo_dict = {
|
if live_options.encoding:
|
||||||
k: v
|
encoding = live_options.encoding
|
||||||
for k, v in lo_dict.items()
|
if live_options.model:
|
||||||
if (not isinstance(v, str) or v != "None") and k != "sample_rate"
|
default_settings.model = live_options.model
|
||||||
}
|
if live_options.language:
|
||||||
if "model" in lo_dict:
|
lang = live_options.language
|
||||||
default_settings.model = lo_dict["model"]
|
default_settings.language = lang.value if isinstance(lang, Language) else lang
|
||||||
if "language" in lo_dict:
|
|
||||||
default_settings.language = lo_dict["language"]
|
|
||||||
if "encoding" in lo_dict:
|
|
||||||
default_settings.encoding = lo_dict["encoding"]
|
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# 3. Apply settings delta (canonical API, always wins)
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
default_settings.apply_update(settings)
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
@@ -221,6 +217,9 @@ class CartesiaSTTService(WebsocketSTTService):
|
|||||||
self._base_url = base_url or "api.cartesia.ai"
|
self._base_url = base_url or "api.cartesia.ai"
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
|
# Init-only audio config (not runtime-updatable).
|
||||||
|
self._encoding = encoding
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
"""Check if the service can generate processing metrics.
|
"""Check if the service can generate processing metrics.
|
||||||
|
|
||||||
@@ -339,7 +338,7 @@ class CartesiaSTTService(WebsocketSTTService):
|
|||||||
params = {
|
params = {
|
||||||
"model": self._settings.model,
|
"model": self._settings.model,
|
||||||
"language": self._settings.language,
|
"language": self._settings.language,
|
||||||
"encoding": self._settings.encoding,
|
"encoding": self._encoding,
|
||||||
"sample_rate": str(self.sample_rate),
|
"sample_rate": str(self.sample_rate),
|
||||||
}
|
}
|
||||||
ws_url = f"wss://{self._base_url}/stt/websocket?{urllib.parse.urlencode(params)}"
|
ws_url = f"wss://{self._base_url}/stt/websocket?{urllib.parse.urlencode(params)}"
|
||||||
|
|||||||
@@ -81,20 +81,16 @@ class DeepgramFluxSTTSettings(STTSettings):
|
|||||||
eot_timeout_ms: Time in ms after speech to finish a turn regardless of EOT
|
eot_timeout_ms: Time in ms after speech to finish a turn regardless of EOT
|
||||||
confidence (default 5000).
|
confidence (default 5000).
|
||||||
keyterm: Keyterms to boost recognition accuracy for specialized terminology.
|
keyterm: Keyterms to boost recognition accuracy for specialized terminology.
|
||||||
mip_opt_out: Opt out of the Deepgram Model Improvement Program (default False).
|
|
||||||
tag: Tags to label requests for identification during usage reporting.
|
tag: Tags to label requests for identification during usage reporting.
|
||||||
min_confidence: Minimum confidence required to create a TranscriptionFrame.
|
min_confidence: Minimum confidence required to create a TranscriptionFrame.
|
||||||
encoding: Audio encoding format (e.g. ``"linear16"``).
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
eager_eot_threshold: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
eager_eot_threshold: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
eot_threshold: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
eot_threshold: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
eot_timeout_ms: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
eot_timeout_ms: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
keyterm: list | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
keyterm: list | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
mip_opt_out: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
tag: list | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
tag: list | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
min_confidence: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
min_confidence: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
encoding: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
|
|
||||||
|
|
||||||
class DeepgramFluxSTTService(WebsocketSTTService):
|
class DeepgramFluxSTTService(WebsocketSTTService):
|
||||||
@@ -158,6 +154,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
api_key: str,
|
api_key: str,
|
||||||
url: str = "wss://api.deepgram.com/v2/listen",
|
url: str = "wss://api.deepgram.com/v2/listen",
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
|
mip_opt_out: Optional[bool] = None,
|
||||||
model: Optional[str] = None,
|
model: Optional[str] = None,
|
||||||
flux_encoding: str = "linear16",
|
flux_encoding: str = "linear16",
|
||||||
params: Optional[InputParams] = None,
|
params: Optional[InputParams] = None,
|
||||||
@@ -170,7 +167,9 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
Args:
|
Args:
|
||||||
api_key: Deepgram API key for authentication. Required for API access.
|
api_key: Deepgram API key for authentication. Required for API access.
|
||||||
url: WebSocket URL for the Deepgram Flux API. Defaults to the preview endpoint.
|
url: WebSocket URL for the Deepgram Flux API. Defaults to the preview endpoint.
|
||||||
sample_rate: Audio sample rate in Hz. If None, uses the rate from params or 16000.
|
sample_rate: Audio sample rate in Hz. If None, uses the pipeline
|
||||||
|
sample rate.
|
||||||
|
mip_opt_out: Opt out of the Deepgram Model Improvement Program.
|
||||||
model: Deepgram Flux model to use for transcription.
|
model: Deepgram Flux model to use for transcription.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
@@ -221,12 +220,10 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
default_settings = DeepgramFluxSTTSettings(
|
default_settings = DeepgramFluxSTTSettings(
|
||||||
model="flux-general-en",
|
model="flux-general-en",
|
||||||
language=Language.EN,
|
language=Language.EN,
|
||||||
encoding=flux_encoding,
|
|
||||||
eager_eot_threshold=None,
|
eager_eot_threshold=None,
|
||||||
eot_threshold=None,
|
eot_threshold=None,
|
||||||
eot_timeout_ms=None,
|
eot_timeout_ms=None,
|
||||||
keyterm=[],
|
keyterm=[],
|
||||||
mip_opt_out=None,
|
|
||||||
tag=[],
|
tag=[],
|
||||||
min_confidence=None,
|
min_confidence=None,
|
||||||
)
|
)
|
||||||
@@ -244,9 +241,10 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
default_settings.eot_threshold = params.eot_threshold
|
default_settings.eot_threshold = params.eot_threshold
|
||||||
default_settings.eot_timeout_ms = params.eot_timeout_ms
|
default_settings.eot_timeout_ms = params.eot_timeout_ms
|
||||||
default_settings.keyterm = params.keyterm or []
|
default_settings.keyterm = params.keyterm or []
|
||||||
default_settings.mip_opt_out = params.mip_opt_out
|
|
||||||
default_settings.tag = params.tag or []
|
default_settings.tag = params.tag or []
|
||||||
default_settings.min_confidence = params.min_confidence
|
default_settings.min_confidence = params.min_confidence
|
||||||
|
if params.mip_opt_out is not None:
|
||||||
|
mip_opt_out = params.mip_opt_out
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# 4. Apply settings delta (canonical API, always wins)
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
@@ -261,8 +259,11 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._url = url
|
self._url = url
|
||||||
self._should_interrupt = should_interrupt
|
self._should_interrupt = should_interrupt
|
||||||
|
self._encoding = flux_encoding
|
||||||
|
self._mip_opt_out = mip_opt_out
|
||||||
self._websocket_url = None
|
self._websocket_url = None
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
# Flux event handlers
|
# Flux event handlers
|
||||||
self._register_event_handler("on_start_of_turn")
|
self._register_event_handler("on_start_of_turn")
|
||||||
self._register_event_handler("on_turn_resumed")
|
self._register_event_handler("on_turn_resumed")
|
||||||
@@ -448,7 +449,7 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
url_params = [
|
url_params = [
|
||||||
f"model={self._settings.model}",
|
f"model={self._settings.model}",
|
||||||
f"sample_rate={self.sample_rate}",
|
f"sample_rate={self.sample_rate}",
|
||||||
f"encoding={self._settings.encoding}",
|
f"encoding={self._encoding}",
|
||||||
]
|
]
|
||||||
|
|
||||||
if self._settings.eager_eot_threshold is not None:
|
if self._settings.eager_eot_threshold is not None:
|
||||||
@@ -460,8 +461,8 @@ class DeepgramFluxSTTService(WebsocketSTTService):
|
|||||||
if self._settings.eot_timeout_ms is not None:
|
if self._settings.eot_timeout_ms is not None:
|
||||||
url_params.append(f"eot_timeout_ms={self._settings.eot_timeout_ms}")
|
url_params.append(f"eot_timeout_ms={self._settings.eot_timeout_ms}")
|
||||||
|
|
||||||
if self._settings.mip_opt_out is not None:
|
if self._mip_opt_out is not None:
|
||||||
url_params.append(f"mip_opt_out={str(self._settings.mip_opt_out).lower()}")
|
url_params.append(f"mip_opt_out={str(self._mip_opt_out).lower()}")
|
||||||
|
|
||||||
# Add keyterm parameters (can have multiple)
|
# Add keyterm parameters (can have multiple)
|
||||||
for keyterm in self._settings.keyterm:
|
for keyterm in self._settings.keyterm:
|
||||||
|
|||||||
@@ -14,7 +14,7 @@ languages, and various Deepgram features.
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
import json
|
import json
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, fields
|
||||||
from typing import Any, AsyncGenerator, Optional
|
from typing import Any, AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -32,32 +32,23 @@ from pipecat.frames.frames import (
|
|||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.aws.sagemaker.bidi_client import SageMakerBidiClient
|
from pipecat.services.aws.sagemaker.bidi_client import SageMakerBidiClient
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.deepgram.stt import DeepgramSTTSettings, LiveOptions
|
||||||
|
from pipecat.services.settings import STTSettings, _warn_deprecated_param, is_given
|
||||||
from pipecat.services.stt_latency import DEEPGRAM_SAGEMAKER_TTFS_P99
|
from pipecat.services.stt_latency import DEEPGRAM_SAGEMAKER_TTFS_P99
|
||||||
from pipecat.services.stt_service import STTService
|
from pipecat.services.stt_service import STTService
|
||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
from pipecat.utils.tracing.service_decorators import traced_stt
|
from pipecat.utils.tracing.service_decorators import traced_stt
|
||||||
|
|
||||||
try:
|
|
||||||
from deepgram import LiveOptions
|
|
||||||
except ModuleNotFoundError as e:
|
|
||||||
logger.error(f"Exception: {e}")
|
|
||||||
logger.error(
|
|
||||||
"In order to use DeepgramSageMakerSTTService, you need to `pip install pipecat-ai[deepgram,sagemaker]`."
|
|
||||||
)
|
|
||||||
raise Exception(f"Missing module: {e}")
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DeepgramSageMakerSTTSettings(STTSettings):
|
class DeepgramSageMakerSTTSettings(DeepgramSTTSettings):
|
||||||
"""Settings for the Deepgram SageMaker STT service.
|
"""Settings for the Deepgram SageMaker STT service.
|
||||||
|
|
||||||
Parameters:
|
Inherits all fields from :class:`DeepgramSTTSettings`.
|
||||||
live_options: Deepgram LiveOptions for the SageMaker connection.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
live_options: Any | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
pass
|
||||||
|
|
||||||
|
|
||||||
class DeepgramSageMakerSTTService(STTService):
|
class DeepgramSageMakerSTTService(STTService):
|
||||||
@@ -72,14 +63,13 @@ class DeepgramSageMakerSTTService(STTService):
|
|||||||
|
|
||||||
- AWS credentials configured (via environment variables, AWS CLI, or instance metadata)
|
- AWS credentials configured (via environment variables, AWS CLI, or instance metadata)
|
||||||
- A deployed SageMaker endpoint with Deepgram model: https://developers.deepgram.com/docs/deploy-amazon-sagemaker
|
- A deployed SageMaker endpoint with Deepgram model: https://developers.deepgram.com/docs/deploy-amazon-sagemaker
|
||||||
- Deepgram SDK for LiveOptions configuration
|
|
||||||
|
|
||||||
Example::
|
Example::
|
||||||
|
|
||||||
stt = DeepgramSageMakerSTTService(
|
stt = DeepgramSageMakerSTTService(
|
||||||
endpoint_name="my-deepgram-endpoint",
|
endpoint_name="my-deepgram-endpoint",
|
||||||
region="us-east-2",
|
region="us-east-2",
|
||||||
live_options=LiveOptions(
|
settings=DeepgramSageMakerSTTSettings(
|
||||||
model="nova-3",
|
model="nova-3",
|
||||||
language="en",
|
language="en",
|
||||||
interim_results=True,
|
interim_results=True,
|
||||||
@@ -95,7 +85,11 @@ class DeepgramSageMakerSTTService(STTService):
|
|||||||
*,
|
*,
|
||||||
endpoint_name: str,
|
endpoint_name: str,
|
||||||
region: str,
|
region: str,
|
||||||
|
encoding: str = "linear16",
|
||||||
|
channels: int = 1,
|
||||||
|
multichannel: bool = False,
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
|
mip_opt_out: Optional[bool] = None,
|
||||||
live_options: Optional[LiveOptions] = None,
|
live_options: Optional[LiveOptions] = None,
|
||||||
settings: Optional[DeepgramSageMakerSTTSettings] = None,
|
settings: Optional[DeepgramSageMakerSTTSettings] = None,
|
||||||
ttfs_p99_latency: Optional[float] = DEEPGRAM_SAGEMAKER_TTFS_P99,
|
ttfs_p99_latency: Optional[float] = DEEPGRAM_SAGEMAKER_TTFS_P99,
|
||||||
@@ -107,11 +101,20 @@ class DeepgramSageMakerSTTService(STTService):
|
|||||||
endpoint_name: Name of the SageMaker endpoint with Deepgram model
|
endpoint_name: Name of the SageMaker endpoint with Deepgram model
|
||||||
deployed (e.g., "my-deepgram-nova-3-endpoint").
|
deployed (e.g., "my-deepgram-nova-3-endpoint").
|
||||||
region: AWS region where the endpoint is deployed (e.g., "us-east-2").
|
region: AWS region where the endpoint is deployed (e.g., "us-east-2").
|
||||||
sample_rate: Audio sample rate in Hz. If None, uses value from
|
encoding: Audio encoding format. Defaults to "linear16".
|
||||||
live_options or defaults to the value from StartFrame.
|
channels: Number of audio channels. Defaults to 1.
|
||||||
live_options: Deepgram LiveOptions configuration. Treated as a
|
multichannel: Transcribe each audio channel independently.
|
||||||
delta from a set of sensible defaults — only the fields you
|
Defaults to False.
|
||||||
set are overridden; all others keep their default values.
|
sample_rate: Audio sample rate in Hz. If None, uses the pipeline
|
||||||
|
sample rate.
|
||||||
|
mip_opt_out: Opt out of Deepgram model improvement program.
|
||||||
|
live_options: Legacy configuration options.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=DeepgramSageMakerSTTSettings(...)`` for
|
||||||
|
runtime-updatable fields and direct init parameters for
|
||||||
|
connection-level config.
|
||||||
|
|
||||||
settings: Runtime-updatable settings. When provided alongside
|
settings: Runtime-updatable settings. When provided alongside
|
||||||
``live_options``, ``settings`` values take precedence (applied
|
``live_options``, ``settings`` values take precedence (applied
|
||||||
after the ``live_options`` merge).
|
after the ``live_options`` merge).
|
||||||
@@ -119,43 +122,63 @@ class DeepgramSageMakerSTTService(STTService):
|
|||||||
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
||||||
**kwargs: Additional arguments passed to the parent STTService.
|
**kwargs: Additional arguments passed to the parent STTService.
|
||||||
"""
|
"""
|
||||||
sample_rate = sample_rate or (live_options.sample_rate if live_options else None)
|
|
||||||
|
|
||||||
default_options = LiveOptions(
|
|
||||||
encoding="linear16",
|
|
||||||
language=Language.EN,
|
|
||||||
model="nova-3",
|
|
||||||
channels=1,
|
|
||||||
interim_results=True,
|
|
||||||
punctuate=True,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 1. Initialize default_settings with hardcoded defaults
|
# 1. Initialize default_settings with hardcoded defaults
|
||||||
default_settings = DeepgramSageMakerSTTSettings(
|
default_settings = DeepgramSageMakerSTTSettings(
|
||||||
model="nova-3",
|
model="nova-3",
|
||||||
language=Language.EN,
|
language=Language.EN,
|
||||||
live_options=default_options,
|
detect_entities=False,
|
||||||
|
diarize=False,
|
||||||
|
dictation=False,
|
||||||
|
endpointing=None,
|
||||||
|
interim_results=True,
|
||||||
|
keyterm=None,
|
||||||
|
keywords=None,
|
||||||
|
numerals=False,
|
||||||
|
profanity_filter=True,
|
||||||
|
punctuate=True,
|
||||||
|
redact=None,
|
||||||
|
replace=None,
|
||||||
|
search=None,
|
||||||
|
smart_format=False,
|
||||||
|
utterance_end_ms=None,
|
||||||
|
vad_events=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. (no deprecated direct args like model= for this service)
|
# 2. Apply live_options overrides — only if settings not provided
|
||||||
|
|
||||||
# 3. Apply live_options overrides — only if settings not provided
|
|
||||||
if live_options is not None:
|
if live_options is not None:
|
||||||
_warn_deprecated_param("live_options", DeepgramSageMakerSTTSettings)
|
_warn_deprecated_param("live_options", DeepgramSageMakerSTTSettings)
|
||||||
if not settings:
|
if not settings:
|
||||||
# Merge user live_options onto defaults
|
# Extract init-only fields from live_options
|
||||||
merged_dict = {**default_options.to_dict(), **live_options.to_dict()}
|
if live_options.sample_rate is not None and sample_rate is None:
|
||||||
merged_live_options = LiveOptions(**merged_dict)
|
sample_rate = live_options.sample_rate
|
||||||
default_settings.live_options = merged_live_options
|
if live_options.encoding is not None:
|
||||||
if hasattr(live_options, "model") and live_options.model is not None:
|
encoding = live_options.encoding
|
||||||
default_settings.model = live_options.model
|
if live_options.channels is not None:
|
||||||
if hasattr(live_options, "language") and live_options.language is not None:
|
channels = live_options.channels
|
||||||
default_settings.language = live_options.language
|
if live_options.multichannel is not None:
|
||||||
|
multichannel = live_options.multichannel
|
||||||
|
if live_options.mip_opt_out is not None:
|
||||||
|
mip_opt_out = live_options.mip_opt_out
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# Build settings delta from remaining fields
|
||||||
|
init_only = {
|
||||||
|
"sample_rate",
|
||||||
|
"encoding",
|
||||||
|
"channels",
|
||||||
|
"multichannel",
|
||||||
|
"mip_opt_out",
|
||||||
|
}
|
||||||
|
lo_dict = {k: v for k, v in live_options.to_dict().items() if k not in init_only}
|
||||||
|
delta = DeepgramSageMakerSTTSettings.from_mapping(lo_dict)
|
||||||
|
default_settings.apply_update(delta)
|
||||||
|
|
||||||
|
# 3. Apply settings delta (canonical API, always wins)
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
default_settings.apply_update(settings)
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
|
# Sync extra to top-level fields so self._settings is unambiguous
|
||||||
|
default_settings._sync_extra_to_fields()
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
sample_rate=sample_rate,
|
sample_rate=sample_rate,
|
||||||
ttfs_p99_latency=ttfs_p99_latency,
|
ttfs_p99_latency=ttfs_p99_latency,
|
||||||
@@ -166,6 +189,12 @@ class DeepgramSageMakerSTTService(STTService):
|
|||||||
self._endpoint_name = endpoint_name
|
self._endpoint_name = endpoint_name
|
||||||
self._region = region
|
self._region = region
|
||||||
|
|
||||||
|
# Init-only connection config (not runtime-updatable).
|
||||||
|
self._encoding = encoding
|
||||||
|
self._channels = channels
|
||||||
|
self._multichannel = multichannel
|
||||||
|
self._mip_opt_out = mip_opt_out
|
||||||
|
|
||||||
self._client: Optional[SageMakerBidiClient] = None
|
self._client: Optional[SageMakerBidiClient] = None
|
||||||
self._response_task: Optional[asyncio.Task] = None
|
self._response_task: Optional[asyncio.Task] = None
|
||||||
self._keepalive_task: Optional[asyncio.Task] = None
|
self._keepalive_task: Optional[asyncio.Task] = None
|
||||||
@@ -185,6 +214,10 @@ class DeepgramSageMakerSTTService(STTService):
|
|||||||
if not changed:
|
if not changed:
|
||||||
return changed
|
return changed
|
||||||
|
|
||||||
|
# Sync extra to fields after the update so self._settings stays unambiguous
|
||||||
|
if isinstance(self._settings, DeepgramSTTSettings):
|
||||||
|
self._settings._sync_extra_to_fields()
|
||||||
|
|
||||||
# TODO: someday we could reconnect here to apply updated settings.
|
# TODO: someday we could reconnect here to apply updated settings.
|
||||||
# Code might look something like the below:
|
# Code might look something like the below:
|
||||||
# await self._disconnect()
|
# await self._disconnect()
|
||||||
@@ -237,6 +270,43 @@ class DeepgramSageMakerSTTService(STTService):
|
|||||||
yield ErrorFrame(error=f"Unknown error occurred: {e}")
|
yield ErrorFrame(error=f"Unknown error occurred: {e}")
|
||||||
yield None
|
yield None
|
||||||
|
|
||||||
|
def _build_query_string(self) -> str:
|
||||||
|
"""Build query string from current settings and init-only connection config."""
|
||||||
|
params = {}
|
||||||
|
s = self._settings
|
||||||
|
|
||||||
|
# Declared Deepgram-specific fields from settings
|
||||||
|
for f in fields(s):
|
||||||
|
if f.name in ("model", "language", "extra") or f.name.startswith("_"):
|
||||||
|
continue
|
||||||
|
value = getattr(s, f.name)
|
||||||
|
if not is_given(value) or value is None:
|
||||||
|
continue
|
||||||
|
params[f.name] = str(value).lower() if isinstance(value, bool) else str(value)
|
||||||
|
|
||||||
|
# model and language
|
||||||
|
if is_given(s.model) and s.model is not None:
|
||||||
|
params["model"] = str(s.model)
|
||||||
|
if is_given(s.language) and s.language is not None:
|
||||||
|
params["language"] = str(s.language)
|
||||||
|
|
||||||
|
# Init-only connection config
|
||||||
|
params["encoding"] = self._encoding
|
||||||
|
params["channels"] = str(self._channels)
|
||||||
|
params["multichannel"] = str(self._multichannel).lower()
|
||||||
|
params["sample_rate"] = str(self.sample_rate)
|
||||||
|
|
||||||
|
if self._mip_opt_out is not None:
|
||||||
|
params["mip_opt_out"] = str(self._mip_opt_out).lower()
|
||||||
|
|
||||||
|
# Any remaining values in extra
|
||||||
|
if s.extra:
|
||||||
|
for key, value in s.extra.items():
|
||||||
|
if value is not None:
|
||||||
|
params[key] = str(value).lower() if isinstance(value, bool) else str(value)
|
||||||
|
|
||||||
|
return "&".join(f"{k}={v}" for k, v in params.items())
|
||||||
|
|
||||||
async def _connect(self):
|
async def _connect(self):
|
||||||
"""Connect to the SageMaker endpoint and start the BiDi session.
|
"""Connect to the SageMaker endpoint and start the BiDi session.
|
||||||
|
|
||||||
@@ -246,21 +316,7 @@ class DeepgramSageMakerSTTService(STTService):
|
|||||||
"""
|
"""
|
||||||
logger.debug("Connecting to Deepgram on SageMaker...")
|
logger.debug("Connecting to Deepgram on SageMaker...")
|
||||||
|
|
||||||
live_options = LiveOptions(
|
query_string = self._build_query_string()
|
||||||
**{**self._settings.live_options.to_dict(), "sample_rate": self.sample_rate}
|
|
||||||
)
|
|
||||||
|
|
||||||
# Build query string from live_options, converting booleans to strings
|
|
||||||
query_params = {}
|
|
||||||
for key, value in live_options.to_dict().items():
|
|
||||||
if value is not None:
|
|
||||||
# Convert boolean values to lowercase strings for Deepgram API
|
|
||||||
if isinstance(value, bool):
|
|
||||||
query_params[key] = str(value).lower()
|
|
||||||
else:
|
|
||||||
query_params[key] = str(value)
|
|
||||||
|
|
||||||
query_string = "&".join(f"{k}={v}" for k, v in query_params.items())
|
|
||||||
|
|
||||||
# Create BiDi client
|
# Create BiDi client
|
||||||
self._client = SageMakerBidiClient(
|
self._client = SageMakerBidiClient(
|
||||||
|
|||||||
@@ -187,8 +187,7 @@ class DeepgramSageMakerTTSService(TTSService):
|
|||||||
logger.debug("Connecting to Deepgram TTS on SageMaker...")
|
logger.debug("Connecting to Deepgram TTS on SageMaker...")
|
||||||
|
|
||||||
query_string = (
|
query_string = (
|
||||||
f"model={self._settings.voice}&encoding={self._settings.encoding}"
|
f"model={self._settings.voice}&encoding={self._encoding}&sample_rate={self.sample_rate}"
|
||||||
f"&sample_rate={self.sample_rate}"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
self._client = SageMakerBidiClient(
|
self._client = SageMakerBidiClient(
|
||||||
|
|||||||
@@ -8,7 +8,7 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from dataclasses import dataclass, field, fields
|
from dataclasses import dataclass, field, fields
|
||||||
from typing import Any, AsyncGenerator, Dict, Optional
|
from typing import Any, AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -26,7 +26,6 @@ from pipecat.frames.frames import (
|
|||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.settings import (
|
from pipecat.services.settings import (
|
||||||
_S,
|
|
||||||
NOT_GIVEN,
|
NOT_GIVEN,
|
||||||
STTSettings,
|
STTSettings,
|
||||||
_NotGiven,
|
_NotGiven,
|
||||||
@@ -57,8 +56,11 @@ class LiveOptions:
|
|||||||
"""Deepgram live transcription options.
|
"""Deepgram live transcription options.
|
||||||
|
|
||||||
Compatibility wrapper that mirrors the ``LiveOptions`` class removed in
|
Compatibility wrapper that mirrors the ``LiveOptions`` class removed in
|
||||||
deepgram-sdk v6. Pass this to :class:`DeepgramSTTService` via the
|
deepgram-sdk v6.
|
||||||
``live_options`` constructor argument.
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=DeepgramSTTSettings(...)`` for runtime-updatable fields
|
||||||
|
and direct ``__init__`` parameters for connection-level config instead.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -179,29 +181,42 @@ class DeepgramSTTSettings(STTSettings):
|
|||||||
|
|
||||||
``model`` and ``language`` are inherited from ``STTSettings`` /
|
``model`` and ``language`` are inherited from ``STTSettings`` /
|
||||||
``ServiceSettings``. Additional Deepgram connection params may
|
``ServiceSettings``. Additional Deepgram connection params may
|
||||||
be passed in through extra ``extra`` (also inherited).
|
be passed in through ``extra`` (also inherited).
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
channels: Number of audio channels.
|
detect_entities: Enable named entity detection.
|
||||||
diarize: Enable speaker diarization.
|
diarize: Enable speaker diarization.
|
||||||
encoding: Audio encoding (e.g. ``"linear16"``).
|
dictation: Enable dictation mode (converts commands to punctuation).
|
||||||
endpointing: Endpointing sensitivity in ms, or ``False`` to disable.
|
endpointing: Endpointing sensitivity in ms, or ``False`` to disable.
|
||||||
interim_results: Whether to emit interim transcriptions.
|
interim_results: Whether to emit interim transcriptions.
|
||||||
|
keyterm: Keyterms to boost (str or list of str).
|
||||||
|
keywords: Keywords to boost (str or list of str).
|
||||||
|
numerals: Convert spoken numbers to numerals.
|
||||||
profanity_filter: Filter profanity from transcripts.
|
profanity_filter: Filter profanity from transcripts.
|
||||||
punctuate: Add punctuation to transcripts.
|
punctuate: Add punctuation to transcripts.
|
||||||
|
redact: Redact sensitive information (str or list of redaction types).
|
||||||
|
replace: Word replacement rules (str or list).
|
||||||
|
search: Search terms to highlight (str or list of str).
|
||||||
smart_format: Apply smart formatting to transcripts.
|
smart_format: Apply smart formatting to transcripts.
|
||||||
|
utterance_end_ms: Silence duration in ms before an utterance-end event.
|
||||||
vad_events: Enable Deepgram VAD speech-started / utterance-end events.
|
vad_events: Enable Deepgram VAD speech-started / utterance-end events.
|
||||||
extra: Additional Deepgram query parameters not covered by the fields above.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
channels: int | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
detect_entities: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
diarize: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
diarize: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
encoding: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
dictation: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
endpointing: Any | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
endpointing: Any | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
interim_results: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
interim_results: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
keyterm: Any | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
keywords: Any | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
numerals: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
profanity_filter: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
profanity_filter: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
punctuate: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
punctuate: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
redact: Any | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
replace: Any | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
search: Any | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
smart_format: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
smart_format: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
utterance_end_ms: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
vad_events: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
vad_events: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
|
||||||
def _sync_extra_to_fields(self) -> None:
|
def _sync_extra_to_fields(self) -> None:
|
||||||
@@ -259,9 +274,16 @@ class DeepgramSTTService(STTService):
|
|||||||
api_key: str,
|
api_key: str,
|
||||||
url: str = "",
|
url: str = "",
|
||||||
base_url: str = "",
|
base_url: str = "",
|
||||||
|
encoding: str = "linear16",
|
||||||
|
channels: int = 1,
|
||||||
|
multichannel: bool = False,
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
|
callback: Optional[str] = None,
|
||||||
|
callback_method: Optional[str] = None,
|
||||||
|
tag: Optional[Any] = None,
|
||||||
|
mip_opt_out: Optional[bool] = None,
|
||||||
live_options: Optional[LiveOptions] = None,
|
live_options: Optional[LiveOptions] = None,
|
||||||
addons: Optional[Dict] = None,
|
addons: Optional[dict] = None,
|
||||||
should_interrupt: bool = True,
|
should_interrupt: bool = True,
|
||||||
settings: Optional[DeepgramSTTSettings] = None,
|
settings: Optional[DeepgramSTTSettings] = None,
|
||||||
ttfs_p99_latency: Optional[float] = DEEPGRAM_TTFS_P99,
|
ttfs_p99_latency: Optional[float] = DEEPGRAM_TTFS_P99,
|
||||||
@@ -277,12 +299,25 @@ class DeepgramSTTService(STTService):
|
|||||||
Parameter `url` is deprecated, use `base_url` instead.
|
Parameter `url` is deprecated, use `base_url` instead.
|
||||||
|
|
||||||
base_url: Custom Deepgram API base URL.
|
base_url: Custom Deepgram API base URL.
|
||||||
sample_rate: Audio sample rate. If None, uses default or live_options value.
|
encoding: Audio encoding format. Defaults to "linear16".
|
||||||
live_options: :class: LiveOptions configuration. Treated as a
|
channels: Number of audio channels. Defaults to 1.
|
||||||
delta from a set of sensible defaults — only the fields you
|
multichannel: Transcribe each audio channel independently.
|
||||||
set are overridden; all others keep their default values.
|
Defaults to False.
|
||||||
|
sample_rate: Audio sample rate in Hz. If None, uses the pipeline
|
||||||
|
sample rate.
|
||||||
|
callback: Callback URL for async transcription delivery.
|
||||||
|
callback_method: HTTP method for the callback (``"GET"`` or ``"POST"``).
|
||||||
|
tag: Custom billing tag.
|
||||||
|
mip_opt_out: Opt out of Deepgram model improvement program.
|
||||||
|
live_options: Legacy configuration options.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=DeepgramSTTSettings(...)`` for runtime-updatable
|
||||||
|
fields and direct init parameters for connection-level config.
|
||||||
|
|
||||||
addons: Additional Deepgram features to enable.
|
addons: Additional Deepgram features to enable.
|
||||||
should_interrupt: Determine whether the bot should be interrupted when Deepgram VAD events are enabled and the system detects that the user is speaking.
|
should_interrupt: Whether to interrupt the bot when Deepgram VAD
|
||||||
|
detects the user is speaking.
|
||||||
|
|
||||||
.. deprecated:: 0.0.99
|
.. deprecated:: 0.0.99
|
||||||
This parameter will be removed along with `vad_events` support.
|
This parameter will be removed along with `vad_events` support.
|
||||||
@@ -297,8 +332,6 @@ class DeepgramSTTService(STTService):
|
|||||||
Note:
|
Note:
|
||||||
The `vad_events` option in LiveOptions is deprecated as of version 0.0.99 and will be removed in a future version. Please use the Silero VAD instead.
|
The `vad_events` option in LiveOptions is deprecated as of version 0.0.99 and will be removed in a future version. Please use the Silero VAD instead.
|
||||||
"""
|
"""
|
||||||
sample_rate = sample_rate or (live_options.sample_rate if live_options else None)
|
|
||||||
|
|
||||||
if url:
|
if url:
|
||||||
import warnings
|
import warnings
|
||||||
|
|
||||||
@@ -314,30 +347,62 @@ class DeepgramSTTService(STTService):
|
|||||||
default_settings = DeepgramSTTSettings(
|
default_settings = DeepgramSTTSettings(
|
||||||
model="nova-3-general",
|
model="nova-3-general",
|
||||||
language=Language.EN,
|
language=Language.EN,
|
||||||
encoding="linear16",
|
detect_entities=False,
|
||||||
channels=1,
|
|
||||||
interim_results=True,
|
|
||||||
smart_format=False,
|
|
||||||
punctuate=True,
|
|
||||||
profanity_filter=True,
|
|
||||||
vad_events=False,
|
|
||||||
diarize=False,
|
diarize=False,
|
||||||
|
dictation=False,
|
||||||
endpointing=None,
|
endpointing=None,
|
||||||
|
interim_results=True,
|
||||||
|
keyterm=None,
|
||||||
|
keywords=None,
|
||||||
|
numerals=False,
|
||||||
|
profanity_filter=True,
|
||||||
|
punctuate=True,
|
||||||
|
redact=None,
|
||||||
|
replace=None,
|
||||||
|
search=None,
|
||||||
|
smart_format=False,
|
||||||
|
utterance_end_ms=None,
|
||||||
|
vad_events=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. (no deprecated direct args like model= for this service)
|
# 2. Apply live_options overrides — only if settings not provided
|
||||||
|
|
||||||
# 3. Apply live_options overrides — only if settings not provided
|
|
||||||
if live_options is not None:
|
if live_options is not None:
|
||||||
_warn_deprecated_param("live_options", DeepgramSTTSettings)
|
_warn_deprecated_param("live_options", DeepgramSTTSettings)
|
||||||
if not settings:
|
if not settings:
|
||||||
lo_dict = live_options.to_dict()
|
# Extract init-only fields from live_options
|
||||||
delta = DeepgramSTTSettings.from_mapping(
|
if live_options.sample_rate is not None and sample_rate is None:
|
||||||
{k: v for k, v in lo_dict.items() if k != "sample_rate"}
|
sample_rate = live_options.sample_rate
|
||||||
)
|
if live_options.encoding is not None:
|
||||||
|
encoding = live_options.encoding
|
||||||
|
if live_options.channels is not None:
|
||||||
|
channels = live_options.channels
|
||||||
|
if live_options.callback is not None:
|
||||||
|
callback = live_options.callback
|
||||||
|
if live_options.callback_method is not None:
|
||||||
|
callback_method = live_options.callback_method
|
||||||
|
if live_options.tag is not None:
|
||||||
|
tag = live_options.tag
|
||||||
|
if live_options.mip_opt_out is not None:
|
||||||
|
mip_opt_out = live_options.mip_opt_out
|
||||||
|
if live_options.multichannel is not None:
|
||||||
|
multichannel = live_options.multichannel
|
||||||
|
|
||||||
|
# Build settings delta from remaining fields
|
||||||
|
init_only = {
|
||||||
|
"sample_rate",
|
||||||
|
"encoding",
|
||||||
|
"channels",
|
||||||
|
"multichannel",
|
||||||
|
"callback",
|
||||||
|
"callback_method",
|
||||||
|
"tag",
|
||||||
|
"mip_opt_out",
|
||||||
|
}
|
||||||
|
lo_dict = {k: v for k, v in live_options.to_dict().items() if k not in init_only}
|
||||||
|
delta = DeepgramSTTSettings.from_mapping(lo_dict)
|
||||||
default_settings.apply_update(delta)
|
default_settings.apply_update(delta)
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# 3. Apply settings delta (canonical API, always wins)
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
default_settings.apply_update(settings)
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
@@ -353,6 +418,13 @@ class DeepgramSTTService(STTService):
|
|||||||
|
|
||||||
self._addons = addons
|
self._addons = addons
|
||||||
self._should_interrupt = should_interrupt
|
self._should_interrupt = should_interrupt
|
||||||
|
self._encoding = encoding
|
||||||
|
self._channels = channels
|
||||||
|
self._multichannel = multichannel
|
||||||
|
self._callback = callback
|
||||||
|
self._callback_method = callback_method
|
||||||
|
self._tag = tag
|
||||||
|
self._mip_opt_out = mip_opt_out
|
||||||
|
|
||||||
if self._settings.vad_events:
|
if self._settings.vad_events:
|
||||||
import warnings
|
import warnings
|
||||||
@@ -487,14 +559,26 @@ class DeepgramSTTService(STTService):
|
|||||||
if is_given(s.language) and s.language is not None:
|
if is_given(s.language) and s.language is not None:
|
||||||
kwargs["language"] = str(s.language)
|
kwargs["language"] = str(s.language)
|
||||||
|
|
||||||
|
# Init-only connection config
|
||||||
|
kwargs["encoding"] = self._encoding
|
||||||
|
kwargs["channels"] = str(self._channels)
|
||||||
|
kwargs["multichannel"] = str(self._multichannel).lower()
|
||||||
|
kwargs["sample_rate"] = str(self.sample_rate)
|
||||||
|
|
||||||
|
if self._callback is not None:
|
||||||
|
kwargs["callback"] = self._callback
|
||||||
|
if self._callback_method is not None:
|
||||||
|
kwargs["callback_method"] = self._callback_method
|
||||||
|
if self._tag is not None:
|
||||||
|
kwargs["tag"] = str(self._tag)
|
||||||
|
if self._mip_opt_out is not None:
|
||||||
|
kwargs["mip_opt_out"] = str(self._mip_opt_out).lower()
|
||||||
|
|
||||||
# Any remaining values in extra (that didn't map to declared fields)
|
# Any remaining values in extra (that didn't map to declared fields)
|
||||||
for key, value in s.extra.items():
|
for key, value in s.extra.items():
|
||||||
if value is not None:
|
if value is not None:
|
||||||
kwargs[key] = str(value).lower() if isinstance(value, bool) else str(value)
|
kwargs[key] = str(value).lower() if isinstance(value, bool) else str(value)
|
||||||
|
|
||||||
# Always inject sample_rate from service level.
|
|
||||||
kwargs["sample_rate"] = str(self.sample_rate)
|
|
||||||
|
|
||||||
if self._addons:
|
if self._addons:
|
||||||
for key, value in self._addons.items():
|
for key, value in self._addons.items():
|
||||||
kwargs[key] = str(value)
|
kwargs[key] = str(value)
|
||||||
|
|||||||
@@ -182,7 +182,8 @@ class ElevenLabsSTTSettings(STTSettings):
|
|||||||
"""Settings for the ElevenLabs file-based STT service.
|
"""Settings for the ElevenLabs file-based STT service.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
tag_audio_events: Whether to include audio event tags in transcription.
|
tag_audio_events: Whether to include audio events like (laughter),
|
||||||
|
(coughing) in the transcription.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
tag_audio_events: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
tag_audio_events: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
@@ -195,7 +196,6 @@ class ElevenLabsRealtimeSTTSettings(STTSettings):
|
|||||||
See ``ElevenLabsRealtimeSTTService.InputParams`` for detailed descriptions.
|
See ``ElevenLabsRealtimeSTTService.InputParams`` for detailed descriptions.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
commit_strategy: How to segment speech - manual (Pipecat VAD) or vad (ElevenLabs VAD).
|
|
||||||
vad_silence_threshold_secs: Seconds of silence before VAD commits (0.3-3.0).
|
vad_silence_threshold_secs: Seconds of silence before VAD commits (0.3-3.0).
|
||||||
vad_threshold: VAD sensitivity (0.1-0.9, lower is more sensitive).
|
vad_threshold: VAD sensitivity (0.1-0.9, lower is more sensitive).
|
||||||
min_speech_duration_ms: Minimum speech duration for VAD (50-2000ms).
|
min_speech_duration_ms: Minimum speech duration for VAD (50-2000ms).
|
||||||
@@ -205,7 +205,6 @@ class ElevenLabsRealtimeSTTSettings(STTSettings):
|
|||||||
include_language_detection: Whether to include language detection in transcripts.
|
include_language_detection: Whether to include language detection in transcripts.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
commit_strategy: CommitStrategy | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
vad_silence_threshold_secs: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
vad_silence_threshold_secs: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
vad_threshold: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
vad_threshold: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
min_speech_duration_ms: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
min_speech_duration_ms: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
@@ -495,6 +494,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
base_url: str = "api.elevenlabs.io",
|
base_url: str = "api.elevenlabs.io",
|
||||||
|
commit_strategy: CommitStrategy = CommitStrategy.MANUAL,
|
||||||
model: Optional[str] = None,
|
model: Optional[str] = None,
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
params: Optional[InputParams] = None,
|
params: Optional[InputParams] = None,
|
||||||
@@ -507,6 +507,9 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
Args:
|
Args:
|
||||||
api_key: ElevenLabs API key for authentication.
|
api_key: ElevenLabs API key for authentication.
|
||||||
base_url: Base URL for ElevenLabs WebSocket API.
|
base_url: Base URL for ElevenLabs WebSocket API.
|
||||||
|
commit_strategy: How to segment speech — ``CommitStrategy.MANUAL``
|
||||||
|
(Pipecat VAD) or ``CommitStrategy.VAD`` (ElevenLabs VAD).
|
||||||
|
Defaults to ``CommitStrategy.MANUAL``.
|
||||||
model: Model ID for transcription.
|
model: Model ID for transcription.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
@@ -528,7 +531,6 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
default_settings = ElevenLabsRealtimeSTTSettings(
|
default_settings = ElevenLabsRealtimeSTTSettings(
|
||||||
model="scribe_v2_realtime",
|
model="scribe_v2_realtime",
|
||||||
language=None,
|
language=None,
|
||||||
commit_strategy=CommitStrategy.MANUAL,
|
|
||||||
vad_silence_threshold_secs=None,
|
vad_silence_threshold_secs=None,
|
||||||
vad_threshold=None,
|
vad_threshold=None,
|
||||||
min_speech_duration_ms=None,
|
min_speech_duration_ms=None,
|
||||||
@@ -548,7 +550,8 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
_warn_deprecated_param("params", ElevenLabsRealtimeSTTSettings)
|
_warn_deprecated_param("params", ElevenLabsRealtimeSTTSettings)
|
||||||
if not settings:
|
if not settings:
|
||||||
default_settings.language = params.language_code
|
default_settings.language = params.language_code
|
||||||
default_settings.commit_strategy = params.commit_strategy
|
if params.commit_strategy != CommitStrategy.MANUAL:
|
||||||
|
commit_strategy = params.commit_strategy
|
||||||
default_settings.vad_silence_threshold_secs = params.vad_silence_threshold_secs
|
default_settings.vad_silence_threshold_secs = params.vad_silence_threshold_secs
|
||||||
default_settings.vad_threshold = params.vad_threshold
|
default_settings.vad_threshold = params.vad_threshold
|
||||||
default_settings.min_speech_duration_ms = params.min_speech_duration_ms
|
default_settings.min_speech_duration_ms = params.min_speech_duration_ms
|
||||||
@@ -575,6 +578,9 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
self._audio_format = "" # initialized in start()
|
self._audio_format = "" # initialized in start()
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
|
# Init-only config (not runtime-updatable).
|
||||||
|
self._commit_strategy = commit_strategy
|
||||||
|
|
||||||
self._connected_event = asyncio.Event()
|
self._connected_event = asyncio.Event()
|
||||||
self._connected_event.set()
|
self._connected_event.set()
|
||||||
|
|
||||||
@@ -651,7 +657,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
await self._start_metrics()
|
await self._start_metrics()
|
||||||
elif isinstance(frame, VADUserStoppedSpeakingFrame):
|
elif isinstance(frame, VADUserStoppedSpeakingFrame):
|
||||||
# Send commit when user stops speaking (manual commit mode)
|
# Send commit when user stops speaking (manual commit mode)
|
||||||
if self._settings.commit_strategy == CommitStrategy.MANUAL:
|
if self._commit_strategy == CommitStrategy.MANUAL:
|
||||||
if self._websocket and self._websocket.state is State.OPEN:
|
if self._websocket and self._websocket.state is State.OPEN:
|
||||||
try:
|
try:
|
||||||
commit_message = {
|
commit_message = {
|
||||||
@@ -754,7 +760,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
params.append(f"language_code={self._settings.language}")
|
params.append(f"language_code={self._settings.language}")
|
||||||
|
|
||||||
params.append(f"audio_format={self._audio_format}")
|
params.append(f"audio_format={self._audio_format}")
|
||||||
params.append(f"commit_strategy={self._settings.commit_strategy.value}")
|
params.append(f"commit_strategy={self._commit_strategy.value}")
|
||||||
|
|
||||||
# Add optional parameters
|
# Add optional parameters
|
||||||
if self._settings.include_timestamps:
|
if self._settings.include_timestamps:
|
||||||
@@ -771,7 +777,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Add VAD parameters if using VAD commit strategy and values are specified
|
# Add VAD parameters if using VAD commit strategy and values are specified
|
||||||
if self._settings.commit_strategy == CommitStrategy.VAD:
|
if self._commit_strategy == CommitStrategy.VAD:
|
||||||
if self._settings.vad_silence_threshold_secs is not None:
|
if self._settings.vad_silence_threshold_secs is not None:
|
||||||
params.append(
|
params.append(
|
||||||
f"vad_silence_threshold_secs={self._settings.vad_silence_threshold_secs}"
|
f"vad_silence_threshold_secs={self._settings.vad_silence_threshold_secs}"
|
||||||
@@ -931,7 +937,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
|
|
||||||
await self._handle_transcription(text, True, language)
|
await self._handle_transcription(text, True, language)
|
||||||
|
|
||||||
finalized = self._settings.commit_strategy == CommitStrategy.MANUAL
|
finalized = self._commit_strategy == CommitStrategy.MANUAL
|
||||||
|
|
||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(
|
TranscriptionFrame(
|
||||||
@@ -975,7 +981,7 @@ class ElevenLabsRealtimeSTTService(WebsocketSTTService):
|
|||||||
|
|
||||||
await self._handle_transcription(text, True, language)
|
await self._handle_transcription(text, True, language)
|
||||||
|
|
||||||
finalized = self._settings.commit_strategy == CommitStrategy.MANUAL
|
finalized = self._commit_strategy == CommitStrategy.MANUAL
|
||||||
|
|
||||||
# This message is sent after committed_transcript when include_timestamps=true.
|
# This message is sent after committed_transcript when include_timestamps=true.
|
||||||
# It contains the full transcript data including text and word-level timestamps.
|
# It contains the full transcript data including text and word-level timestamps.
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ transcription using segmented audio processing.
|
|||||||
|
|
||||||
import base64
|
import base64
|
||||||
import os
|
import os
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
from typing import AsyncGenerator, Optional
|
from typing import AsyncGenerator, Optional
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
@@ -20,7 +20,7 @@ from loguru import logger
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
from pipecat.frames.frames import ErrorFrame, Frame, TranscriptionFrame
|
from pipecat.frames.frames import ErrorFrame, Frame, TranscriptionFrame
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.settings import STTSettings, _warn_deprecated_param
|
||||||
from pipecat.services.stt_latency import FAL_TTFS_P99
|
from pipecat.services.stt_latency import FAL_TTFS_P99
|
||||||
from pipecat.services.stt_service import SegmentedSTTService
|
from pipecat.services.stt_service import SegmentedSTTService
|
||||||
from pipecat.transcriptions.language import Language, resolve_language
|
from pipecat.transcriptions.language import Language, resolve_language
|
||||||
@@ -143,18 +143,9 @@ def language_to_fal_language(language: Language) -> Optional[str]:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class FalSTTSettings(STTSettings):
|
class FalSTTSettings(STTSettings):
|
||||||
"""Settings for the Fal Wizper STT service.
|
"""Settings for the Fal Wizper STT service."""
|
||||||
|
|
||||||
Parameters:
|
pass
|
||||||
task: Task to perform ('transcribe' or 'translate'). Defaults to
|
|
||||||
'transcribe'.
|
|
||||||
chunk_level: Level of chunking ('segment'). Defaults to 'segment'.
|
|
||||||
version: Version of Wizper model to use. Defaults to '3'.
|
|
||||||
"""
|
|
||||||
|
|
||||||
task: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
chunk_level: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
version: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
|
|
||||||
|
|
||||||
class FalSTTService(SegmentedSTTService):
|
class FalSTTService(SegmentedSTTService):
|
||||||
@@ -189,6 +180,9 @@ class FalSTTService(SegmentedSTTService):
|
|||||||
*,
|
*,
|
||||||
api_key: Optional[str] = None,
|
api_key: Optional[str] = None,
|
||||||
aiohttp_session: Optional[aiohttp.ClientSession] = None,
|
aiohttp_session: Optional[aiohttp.ClientSession] = None,
|
||||||
|
task: str = "transcribe",
|
||||||
|
chunk_level: str = "segment",
|
||||||
|
version: str = "3",
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
params: Optional[InputParams] = None,
|
params: Optional[InputParams] = None,
|
||||||
settings: Optional[FalSTTSettings] = None,
|
settings: Optional[FalSTTSettings] = None,
|
||||||
@@ -201,11 +195,16 @@ class FalSTTService(SegmentedSTTService):
|
|||||||
api_key: Fal API key. If not provided, will check FAL_KEY environment variable.
|
api_key: Fal API key. If not provided, will check FAL_KEY environment variable.
|
||||||
aiohttp_session: Optional aiohttp ClientSession for HTTP requests.
|
aiohttp_session: Optional aiohttp ClientSession for HTTP requests.
|
||||||
If not provided, a session will be created and managed internally.
|
If not provided, a session will be created and managed internally.
|
||||||
|
task: Task to perform (``"transcribe"`` or ``"translate"``).
|
||||||
|
Defaults to ``"transcribe"``.
|
||||||
|
chunk_level: Level of chunking (``"segment"``). Defaults to ``"segment"``.
|
||||||
|
version: Version of Wizper model to use. Defaults to ``"3"``.
|
||||||
sample_rate: Audio sample rate in Hz. If not provided, uses the pipeline's rate.
|
sample_rate: Audio sample rate in Hz. If not provided, uses the pipeline's rate.
|
||||||
params: Configuration parameters for the Wizper API.
|
params: Configuration parameters for the Wizper API.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=FalSTTSettings(...)`` instead.
|
Use ``settings=FalSTTSettings(...)`` for model/language and
|
||||||
|
direct init parameters for task/chunk_level/version instead.
|
||||||
|
|
||||||
settings: Runtime-updatable settings. When provided alongside deprecated
|
settings: Runtime-updatable settings. When provided alongside deprecated
|
||||||
parameters, ``settings`` values take precedence.
|
parameters, ``settings`` values take precedence.
|
||||||
@@ -217,9 +216,6 @@ class FalSTTService(SegmentedSTTService):
|
|||||||
default_settings = FalSTTSettings(
|
default_settings = FalSTTSettings(
|
||||||
model=None,
|
model=None,
|
||||||
language=language_to_fal_language(Language.EN) or "en",
|
language=language_to_fal_language(Language.EN) or "en",
|
||||||
task="transcribe",
|
|
||||||
chunk_level="segment",
|
|
||||||
version="3",
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. (no deprecated direct args for this service)
|
# 2. (no deprecated direct args for this service)
|
||||||
@@ -231,9 +227,12 @@ class FalSTTService(SegmentedSTTService):
|
|||||||
default_settings.language = (
|
default_settings.language = (
|
||||||
language_to_fal_language(params.language) if params.language else "en"
|
language_to_fal_language(params.language) if params.language else "en"
|
||||||
)
|
)
|
||||||
default_settings.task = params.task
|
if params.task != "transcribe":
|
||||||
default_settings.chunk_level = params.chunk_level
|
task = params.task
|
||||||
default_settings.version = params.version
|
if params.chunk_level != "segment":
|
||||||
|
chunk_level = params.chunk_level
|
||||||
|
if params.version != "3":
|
||||||
|
version = params.version
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# 4. Apply settings delta (canonical API, always wins)
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
@@ -246,6 +245,10 @@ class FalSTTService(SegmentedSTTService):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
self._task = task
|
||||||
|
self._chunk_level = chunk_level
|
||||||
|
self._version = version
|
||||||
|
|
||||||
self._api_key = api_key or os.getenv("FAL_KEY", "")
|
self._api_key = api_key or os.getenv("FAL_KEY", "")
|
||||||
if not self._api_key:
|
if not self._api_key:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -301,7 +304,15 @@ class FalSTTService(SegmentedSTTService):
|
|||||||
self._session = aiohttp.ClientSession()
|
self._session = aiohttp.ClientSession()
|
||||||
|
|
||||||
data_uri = f"data:audio/x-wav;base64,{base64.b64encode(audio).decode()}"
|
data_uri = f"data:audio/x-wav;base64,{base64.b64encode(audio).decode()}"
|
||||||
payload = {"audio_url": data_uri, **self._settings.given_fields()}
|
payload: dict = {"audio_url": data_uri}
|
||||||
|
if self._settings.language is not None:
|
||||||
|
payload["language"] = self._settings.language
|
||||||
|
if self._task is not None:
|
||||||
|
payload["task"] = self._task
|
||||||
|
if self._chunk_level is not None:
|
||||||
|
payload["chunk_level"] = self._chunk_level
|
||||||
|
if self._version is not None:
|
||||||
|
payload["version"] = self._version
|
||||||
headers = {
|
headers = {
|
||||||
"Authorization": f"Key {self._api_key}",
|
"Authorization": f"Key {self._api_key}",
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
|
|||||||
@@ -152,6 +152,10 @@ class MessagesConfig(BaseModel):
|
|||||||
class GladiaInputParams(BaseModel):
|
class GladiaInputParams(BaseModel):
|
||||||
"""Configuration parameters for the Gladia STT service.
|
"""Configuration parameters for the Gladia STT service.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=GladiaSTTSettings(...)`` for runtime-updatable
|
||||||
|
fields and direct init parameters for encoding/bit_depth/channels.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
encoding: Audio encoding format
|
encoding: Audio encoding format
|
||||||
bit_depth: Audio bit depth
|
bit_depth: Audio bit depth
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ import base64
|
|||||||
import json
|
import json
|
||||||
import warnings
|
import warnings
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any, AsyncGenerator, Dict, Literal, Optional
|
from typing import Any, AsyncGenerator, Literal, Optional
|
||||||
|
|
||||||
import aiohttp
|
import aiohttp
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -191,28 +191,22 @@ class GladiaSTTSettings(STTSettings):
|
|||||||
"""Settings for Gladia STT service.
|
"""Settings for Gladia STT service.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
encoding: Audio encoding format.
|
language_config: Language detection and handling configuration.
|
||||||
bit_depth: Audio bit depth.
|
|
||||||
channels: Number of audio channels.
|
|
||||||
custom_metadata: Additional metadata to include with requests.
|
custom_metadata: Additional metadata to include with requests.
|
||||||
endpointing: Silence duration in seconds to mark end of speech.
|
endpointing: Silence duration in seconds to mark end of speech.
|
||||||
maximum_duration_without_endpointing: Maximum utterance duration without silence.
|
maximum_duration_without_endpointing: Maximum utterance duration without silence.
|
||||||
language_config: Detailed language configuration.
|
|
||||||
pre_processing: Audio pre-processing options.
|
pre_processing: Audio pre-processing options.
|
||||||
realtime_processing: Real-time processing features.
|
realtime_processing: Real-time processing features.
|
||||||
messages_config: WebSocket message filtering options.
|
messages_config: WebSocket message filtering options.
|
||||||
enable_vad: Enable VAD to trigger end of utterance detection.
|
enable_vad: Enable VAD to trigger end of utterance detection.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
encoding: str | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
language_config: LanguageConfig | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
bit_depth: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
custom_metadata: dict[str, Any] | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
channels: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
custom_metadata: Dict[str, Any] | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
endpointing: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
endpointing: float | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
maximum_duration_without_endpointing: int | None | _NotGiven = field(
|
maximum_duration_without_endpointing: int | None | _NotGiven = field(
|
||||||
default_factory=lambda: NOT_GIVEN
|
default_factory=lambda: NOT_GIVEN
|
||||||
)
|
)
|
||||||
language_config: LanguageConfig | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
pre_processing: PreProcessingConfig | None | _NotGiven = field(
|
pre_processing: PreProcessingConfig | None | _NotGiven = field(
|
||||||
default_factory=lambda: NOT_GIVEN
|
default_factory=lambda: NOT_GIVEN
|
||||||
)
|
)
|
||||||
@@ -247,6 +241,9 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
api_key: str,
|
api_key: str,
|
||||||
region: Literal["us-west", "eu-west"] | None = None,
|
region: Literal["us-west", "eu-west"] | None = None,
|
||||||
url: str = "https://api.gladia.io/v2/live",
|
url: str = "https://api.gladia.io/v2/live",
|
||||||
|
encoding: str = "wav/pcm",
|
||||||
|
bit_depth: int = 16,
|
||||||
|
channels: int = 1,
|
||||||
confidence: Optional[float] = None,
|
confidence: Optional[float] = None,
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
model: Optional[str] = None,
|
model: Optional[str] = None,
|
||||||
@@ -263,6 +260,9 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
api_key: Gladia API key for authentication.
|
api_key: Gladia API key for authentication.
|
||||||
region: Region used to process audio. eu-west or us-west. Defaults to eu-west.
|
region: Region used to process audio. eu-west or us-west. Defaults to eu-west.
|
||||||
url: Gladia API URL. Defaults to "https://api.gladia.io/v2/live".
|
url: Gladia API URL. Defaults to "https://api.gladia.io/v2/live".
|
||||||
|
encoding: Audio encoding format. Defaults to ``"wav/pcm"``.
|
||||||
|
bit_depth: Audio bit depth. Defaults to 16.
|
||||||
|
channels: Number of audio channels. Defaults to 1.
|
||||||
confidence: Minimum confidence threshold for transcriptions (0.0-1.0).
|
confidence: Minimum confidence threshold for transcriptions (0.0-1.0).
|
||||||
|
|
||||||
.. deprecated:: 0.0.86
|
.. deprecated:: 0.0.86
|
||||||
@@ -278,7 +278,8 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
params: Additional configuration parameters for Gladia service.
|
params: Additional configuration parameters for Gladia service.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=GladiaSTTSettings(...)`` instead.
|
Use ``settings=GladiaSTTSettings(...)`` for runtime-updatable
|
||||||
|
fields and direct init parameters for encoding/bit_depth/channels.
|
||||||
|
|
||||||
max_buffer_size: Maximum size of audio buffer in bytes. Defaults to 20MB.
|
max_buffer_size: Maximum size of audio buffer in bytes. Defaults to 20MB.
|
||||||
should_interrupt: Determine whether the bot should be interrupted when
|
should_interrupt: Determine whether the bot should be interrupted when
|
||||||
@@ -303,13 +304,10 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
default_settings = GladiaSTTSettings(
|
default_settings = GladiaSTTSettings(
|
||||||
model="solaria-1",
|
model="solaria-1",
|
||||||
language=None,
|
language=None,
|
||||||
encoding="wav/pcm",
|
language_config=None,
|
||||||
bit_depth=16,
|
|
||||||
channels=1,
|
|
||||||
custom_metadata=None,
|
custom_metadata=None,
|
||||||
endpointing=None,
|
endpointing=None,
|
||||||
maximum_duration_without_endpointing=5,
|
maximum_duration_without_endpointing=5,
|
||||||
language_config=None,
|
|
||||||
pre_processing=None,
|
pre_processing=None,
|
||||||
realtime_processing=None,
|
realtime_processing=None,
|
||||||
messages_config=None,
|
messages_config=None,
|
||||||
@@ -334,9 +332,13 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
stacklevel=2,
|
stacklevel=2,
|
||||||
)
|
)
|
||||||
if not settings:
|
if not settings:
|
||||||
default_settings.encoding = params.encoding
|
# Extract init-only fields from params
|
||||||
default_settings.bit_depth = params.bit_depth
|
if params.encoding is not None:
|
||||||
default_settings.channels = params.channels
|
encoding = params.encoding
|
||||||
|
if params.bit_depth is not None:
|
||||||
|
bit_depth = params.bit_depth
|
||||||
|
if params.channels is not None:
|
||||||
|
channels = params.channels
|
||||||
default_settings.custom_metadata = params.custom_metadata
|
default_settings.custom_metadata = params.custom_metadata
|
||||||
default_settings.endpointing = params.endpointing
|
default_settings.endpointing = params.endpointing
|
||||||
default_settings.maximum_duration_without_endpointing = (
|
default_settings.maximum_duration_without_endpointing = (
|
||||||
@@ -347,14 +349,14 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
default_settings.messages_config = params.messages_config
|
default_settings.messages_config = params.messages_config
|
||||||
default_settings.enable_vad = params.enable_vad
|
default_settings.enable_vad = params.enable_vad
|
||||||
# Resolve deprecated language → language_config at init time
|
# Resolve deprecated language → language_config at init time
|
||||||
language_config = params.language_config
|
if params.language_config:
|
||||||
if not language_config and params.language:
|
default_settings.language_config = params.language_config
|
||||||
|
elif params.language:
|
||||||
language_code = self.language_to_service_language(params.language)
|
language_code = self.language_to_service_language(params.language)
|
||||||
if language_code:
|
if language_code:
|
||||||
language_config = LanguageConfig(
|
default_settings.language_config = LanguageConfig(
|
||||||
languages=[language_code], code_switching=False
|
languages=[language_code], code_switching=False
|
||||||
)
|
)
|
||||||
default_settings.language_config = language_config
|
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# 4. Apply settings delta (canonical API, always wins)
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
@@ -374,6 +376,11 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
self._url = url
|
self._url = url
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
|
# Init-only connection config
|
||||||
|
self._encoding = encoding
|
||||||
|
self._bit_depth = bit_depth
|
||||||
|
self._channels = channels
|
||||||
|
|
||||||
# Session management
|
# Session management
|
||||||
self._session_url = None
|
self._session_url = None
|
||||||
self._session_id = None
|
self._session_id = None
|
||||||
@@ -411,14 +418,14 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
"""
|
"""
|
||||||
return language_to_gladia_language(language)
|
return language_to_gladia_language(language)
|
||||||
|
|
||||||
def _prepare_settings(self) -> Dict[str, Any]:
|
def _prepare_settings(self) -> dict[str, Any]:
|
||||||
s = self._settings
|
s = self._settings
|
||||||
|
|
||||||
settings = {
|
settings = {
|
||||||
"encoding": s.encoding or "wav/pcm",
|
"encoding": self._encoding or "wav/pcm",
|
||||||
"bit_depth": s.bit_depth or 16,
|
"bit_depth": self._bit_depth or 16,
|
||||||
"sample_rate": self.sample_rate,
|
"sample_rate": self.sample_rate,
|
||||||
"channels": s.channels or 1,
|
"channels": self._channels or 1,
|
||||||
"model": s.model,
|
"model": s.model,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -610,7 +617,7 @@ class GladiaSTTService(WebsocketSTTService):
|
|||||||
self._websocket = None
|
self._websocket = None
|
||||||
await self._call_event_handler("on_disconnected")
|
await self._call_event_handler("on_disconnected")
|
||||||
|
|
||||||
async def _setup_gladia(self, settings: Dict[str, Any]):
|
async def _setup_gladia(self, settings: dict[str, Any]):
|
||||||
async with aiohttp.ClientSession() as session:
|
async with aiohttp.ClientSession() as session:
|
||||||
params = {}
|
params = {}
|
||||||
if self._region:
|
if self._region:
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ WebSocket API for streaming audio transcription.
|
|||||||
|
|
||||||
import base64
|
import base64
|
||||||
import json
|
import json
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, Optional
|
from typing import Any, AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -28,7 +28,7 @@ from pipecat.frames.frames import (
|
|||||||
VADUserStoppedSpeakingFrame,
|
VADUserStoppedSpeakingFrame,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.settings import STTSettings, _warn_deprecated_param
|
||||||
from pipecat.services.stt_latency import GRADIUM_TTFS_P99
|
from pipecat.services.stt_latency import GRADIUM_TTFS_P99
|
||||||
from pipecat.services.stt_service import WebsocketSTTService
|
from pipecat.services.stt_service import WebsocketSTTService
|
||||||
from pipecat.transcriptions.language import Language, resolve_language
|
from pipecat.transcriptions.language import Language, resolve_language
|
||||||
@@ -68,14 +68,9 @@ def language_to_gradium_language(language: Language) -> Optional[str]:
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class GradiumSTTSettings(STTSettings):
|
class GradiumSTTSettings(STTSettings):
|
||||||
"""Settings for the Gradium STT service.
|
"""Settings for the Gradium STT service."""
|
||||||
|
|
||||||
Parameters:
|
pass
|
||||||
delay_in_frames: Delay in audio frames (80ms each) before text is
|
|
||||||
generated. Higher delays allow more context but increase latency.
|
|
||||||
"""
|
|
||||||
|
|
||||||
delay_in_frames: int | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
|
|
||||||
|
|
||||||
class GradiumSTTService(WebsocketSTTService):
|
class GradiumSTTService(WebsocketSTTService):
|
||||||
@@ -112,6 +107,7 @@ class GradiumSTTService(WebsocketSTTService):
|
|||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
api_endpoint_base_url: str = "wss://eu.api.gradium.ai/api/speech/asr",
|
api_endpoint_base_url: str = "wss://eu.api.gradium.ai/api/speech/asr",
|
||||||
|
delay_in_frames: Optional[int] = None,
|
||||||
params: Optional[InputParams] = None,
|
params: Optional[InputParams] = None,
|
||||||
json_config: Optional[str] = None,
|
json_config: Optional[str] = None,
|
||||||
settings: Optional[GradiumSTTSettings] = None,
|
settings: Optional[GradiumSTTSettings] = None,
|
||||||
@@ -123,6 +119,9 @@ class GradiumSTTService(WebsocketSTTService):
|
|||||||
Args:
|
Args:
|
||||||
api_key: Gradium API key for authentication.
|
api_key: Gradium API key for authentication.
|
||||||
api_endpoint_base_url: WebSocket endpoint URL. Defaults to Gradium's streaming endpoint.
|
api_endpoint_base_url: WebSocket endpoint URL. Defaults to Gradium's streaming endpoint.
|
||||||
|
delay_in_frames: Delay in audio frames (80ms each) before text is
|
||||||
|
generated. Higher delays allow more context but increase latency.
|
||||||
|
Allowed values: 7, 8, 10, 12, 14, 16, 20, 24, 36, 48.
|
||||||
params: Configuration parameters for language and delay settings.
|
params: Configuration parameters for language and delay settings.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
@@ -152,7 +151,6 @@ class GradiumSTTService(WebsocketSTTService):
|
|||||||
default_settings = GradiumSTTSettings(
|
default_settings = GradiumSTTSettings(
|
||||||
model=None,
|
model=None,
|
||||||
language=None,
|
language=None,
|
||||||
delay_in_frames=None,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. (no deprecated direct args for this service)
|
# 2. (no deprecated direct args for this service)
|
||||||
@@ -162,7 +160,8 @@ class GradiumSTTService(WebsocketSTTService):
|
|||||||
_warn_deprecated_param("params", GradiumSTTSettings)
|
_warn_deprecated_param("params", GradiumSTTSettings)
|
||||||
if not settings:
|
if not settings:
|
||||||
default_settings.language = params.language
|
default_settings.language = params.language
|
||||||
default_settings.delay_in_frames = params.delay_in_frames
|
if params.delay_in_frames is not None:
|
||||||
|
delay_in_frames = params.delay_in_frames
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# 4. Apply settings delta (canonical API, always wins)
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
@@ -179,6 +178,7 @@ class GradiumSTTService(WebsocketSTTService):
|
|||||||
self._api_endpoint_base_url = api_endpoint_base_url
|
self._api_endpoint_base_url = api_endpoint_base_url
|
||||||
self._websocket = None
|
self._websocket = None
|
||||||
self._json_config = json_config
|
self._json_config = json_config
|
||||||
|
self._config_delay_in_frames = delay_in_frames
|
||||||
|
|
||||||
self._receive_task = None
|
self._receive_task = None
|
||||||
|
|
||||||
@@ -358,8 +358,8 @@ class GradiumSTTService(WebsocketSTTService):
|
|||||||
gradium_language = language_to_gradium_language(self._settings.language)
|
gradium_language = language_to_gradium_language(self._settings.language)
|
||||||
if gradium_language:
|
if gradium_language:
|
||||||
json_config["language"] = gradium_language
|
json_config["language"] = gradium_language
|
||||||
if self._settings.delay_in_frames:
|
if self._config_delay_in_frames:
|
||||||
json_config["delay_in_frames"] = self._settings.delay_in_frames
|
json_config["delay_in_frames"] = self._config_delay_in_frames
|
||||||
if json_config:
|
if json_config:
|
||||||
setup_msg["json_config"] = json_config
|
setup_msg["json_config"] = json_config
|
||||||
await self._websocket.send(json.dumps(setup_msg))
|
await self._websocket.send(json.dumps(setup_msg))
|
||||||
|
|||||||
@@ -6,6 +6,7 @@
|
|||||||
|
|
||||||
"""Groq speech-to-text service implementation using Whisper models."""
|
"""Groq speech-to-text service implementation using Whisper models."""
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from pipecat.services.settings import _warn_deprecated_param
|
from pipecat.services.settings import _warn_deprecated_param
|
||||||
@@ -18,6 +19,17 @@ from pipecat.services.whisper.base_stt import (
|
|||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class GroqSTTSettings(BaseWhisperSTTSettings):
|
||||||
|
"""Settings for the Groq STT service.
|
||||||
|
|
||||||
|
Parameters:
|
||||||
|
prompt: Optional prompt text to guide transcription style.
|
||||||
|
"""
|
||||||
|
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class GroqSTTService(BaseWhisperSTTService):
|
class GroqSTTService(BaseWhisperSTTService):
|
||||||
"""Groq Whisper speech-to-text service.
|
"""Groq Whisper speech-to-text service.
|
||||||
|
|
||||||
@@ -25,6 +37,8 @@ class GroqSTTService(BaseWhisperSTTService):
|
|||||||
set via the api_key parameter or GROQ_API_KEY environment variable.
|
set via the api_key parameter or GROQ_API_KEY environment variable.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_settings: GroqSTTSettings
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -34,7 +48,7 @@ class GroqSTTService(BaseWhisperSTTService):
|
|||||||
language: Optional[Language] = None,
|
language: Optional[Language] = None,
|
||||||
prompt: Optional[str] = None,
|
prompt: Optional[str] = None,
|
||||||
temperature: Optional[float] = None,
|
temperature: Optional[float] = None,
|
||||||
settings: Optional[BaseWhisperSTTSettings] = None,
|
settings: Optional[GroqSTTSettings] = None,
|
||||||
ttfs_p99_latency: Optional[float] = GROQ_TTFS_P99,
|
ttfs_p99_latency: Optional[float] = GROQ_TTFS_P99,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
@@ -44,24 +58,24 @@ class GroqSTTService(BaseWhisperSTTService):
|
|||||||
model: Whisper model to use.
|
model: Whisper model to use.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(model=...)`` instead.
|
Use ``settings=GroqSTTSettings(model=...)`` instead.
|
||||||
|
|
||||||
api_key: Groq API key. Defaults to None.
|
api_key: Groq API key. Defaults to None.
|
||||||
base_url: API base URL. Defaults to "https://api.groq.com/openai/v1".
|
base_url: API base URL. Defaults to "https://api.groq.com/openai/v1".
|
||||||
language: Language of the audio input.
|
language: Language of the audio input.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(language=...)`` instead.
|
Use ``settings=GroqSTTSettings(language=...)`` instead.
|
||||||
|
|
||||||
prompt: Optional text to guide the model's style or continue a previous segment.
|
prompt: Optional text to guide the model's style or continue a previous segment.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(prompt=...)`` instead.
|
Use ``settings=GroqSTTSettings(prompt=...)`` instead.
|
||||||
|
|
||||||
temperature: Optional sampling temperature between 0 and 1.
|
temperature: Optional sampling temperature between 0 and 1.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(temperature=...)`` instead.
|
Use ``settings=GroqSTTSettings(temperature=...)`` instead.
|
||||||
|
|
||||||
settings: Runtime-updatable settings. When provided alongside deprecated
|
settings: Runtime-updatable settings. When provided alongside deprecated
|
||||||
parameters, ``settings`` values take precedence.
|
parameters, ``settings`` values take precedence.
|
||||||
@@ -70,24 +84,25 @@ class GroqSTTService(BaseWhisperSTTService):
|
|||||||
**kwargs: Additional arguments passed to BaseWhisperSTTService.
|
**kwargs: Additional arguments passed to BaseWhisperSTTService.
|
||||||
"""
|
"""
|
||||||
# --- 1. Hardcoded defaults ---
|
# --- 1. Hardcoded defaults ---
|
||||||
default_settings = BaseWhisperSTTSettings(
|
default_settings = GroqSTTSettings(
|
||||||
model="whisper-large-v3-turbo",
|
model="whisper-large-v3-turbo",
|
||||||
language=self.language_to_service_language(Language.EN),
|
language=self.language_to_service_language(Language.EN),
|
||||||
base_url=base_url,
|
prompt=None,
|
||||||
|
temperature=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# --- 2. Deprecated direct-arg overrides ---
|
# --- 2. Deprecated direct-arg overrides ---
|
||||||
if model is not None:
|
if model is not None:
|
||||||
_warn_deprecated_param("model", BaseWhisperSTTSettings, "model")
|
_warn_deprecated_param("model", GroqSTTSettings, "model")
|
||||||
default_settings.model = model
|
default_settings.model = model
|
||||||
if language is not None:
|
if language is not None:
|
||||||
_warn_deprecated_param("language", BaseWhisperSTTSettings, "language")
|
_warn_deprecated_param("language", GroqSTTSettings, "language")
|
||||||
default_settings.language = self.language_to_service_language(language)
|
default_settings.language = self.language_to_service_language(language)
|
||||||
if prompt is not None:
|
if prompt is not None:
|
||||||
_warn_deprecated_param("prompt", BaseWhisperSTTSettings, "prompt")
|
_warn_deprecated_param("prompt", GroqSTTSettings, "prompt")
|
||||||
default_settings.prompt = prompt
|
default_settings.prompt = prompt
|
||||||
if temperature is not None:
|
if temperature is not None:
|
||||||
_warn_deprecated_param("temperature", BaseWhisperSTTSettings, "temperature")
|
_warn_deprecated_param("temperature", GroqSTTSettings, "temperature")
|
||||||
default_settings.temperature = temperature
|
default_settings.temperature = temperature
|
||||||
|
|
||||||
# --- 3. (no params object for this service) ---
|
# --- 3. (no params object for this service) ---
|
||||||
@@ -105,7 +120,7 @@ class GroqSTTService(BaseWhisperSTTService):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def _transcribe(self, audio: bytes) -> Transcription:
|
async def _transcribe(self, audio: bytes) -> Transcription:
|
||||||
assert self._language is not None # Assigned in the BaseWhisperSTTService class
|
assert self._settings.language is not None
|
||||||
|
|
||||||
# Build kwargs dict with only set parameters
|
# Build kwargs dict with only set parameters
|
||||||
kwargs = {
|
kwargs = {
|
||||||
@@ -113,13 +128,13 @@ class GroqSTTService(BaseWhisperSTTService):
|
|||||||
"model": self._settings.model,
|
"model": self._settings.model,
|
||||||
# Use verbose_json to get probability metrics
|
# Use verbose_json to get probability metrics
|
||||||
"response_format": "verbose_json" if self._include_prob_metrics else "json",
|
"response_format": "verbose_json" if self._include_prob_metrics else "json",
|
||||||
"language": self._language,
|
"language": self._settings.language,
|
||||||
}
|
}
|
||||||
|
|
||||||
if self._prompt is not None:
|
if self._settings.prompt is not None:
|
||||||
kwargs["prompt"] = self._prompt
|
kwargs["prompt"] = self._settings.prompt
|
||||||
|
|
||||||
if self._temperature is not None:
|
if self._settings.temperature is not None:
|
||||||
kwargs["temperature"] = self._temperature
|
kwargs["temperature"] = self._settings.temperature
|
||||||
|
|
||||||
return await self._client.audio.transcriptions.create(**kwargs)
|
return await self._client.audio.transcriptions.create(**kwargs)
|
||||||
|
|||||||
@@ -8,7 +8,7 @@
|
|||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from concurrent.futures import CancelledError as FuturesCancelledError
|
from concurrent.futures import CancelledError as FuturesCancelledError
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, List, Mapping, Optional
|
from typing import Any, AsyncGenerator, List, Mapping, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -23,7 +23,7 @@ from pipecat.frames.frames import (
|
|||||||
StartFrame,
|
StartFrame,
|
||||||
TranscriptionFrame,
|
TranscriptionFrame,
|
||||||
)
|
)
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.settings import STTSettings, _warn_deprecated_param
|
||||||
from pipecat.services.stt_latency import NVIDIA_TTFS_P99
|
from pipecat.services.stt_latency import NVIDIA_TTFS_P99
|
||||||
from pipecat.services.stt_service import SegmentedSTTService, STTService
|
from pipecat.services.stt_service import SegmentedSTTService, STTService
|
||||||
from pipecat.transcriptions.language import Language, resolve_language
|
from pipecat.transcriptions.language import Language, resolve_language
|
||||||
@@ -110,11 +110,11 @@ class NvidiaSegmentedSTTSettings(STTSettings):
|
|||||||
boosted_lm_score: Score boost for specified words.
|
boosted_lm_score: Score boost for specified words.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
profanity_filter: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
profanity_filter: bool = False
|
||||||
automatic_punctuation: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
automatic_punctuation: bool = True
|
||||||
verbatim_transcripts: bool | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
verbatim_transcripts: bool = False
|
||||||
boosted_lm_words: List[str] | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
boosted_lm_words: Optional[List[str]] = None
|
||||||
boosted_lm_score: float | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
boosted_lm_score: float = 4.0
|
||||||
|
|
||||||
|
|
||||||
class NvidiaSTTService(STTService):
|
class NvidiaSTTService(STTService):
|
||||||
@@ -586,19 +586,18 @@ class NvidiaSegmentedSTTService(SegmentedSTTService):
|
|||||||
def _create_recognition_config(self):
|
def _create_recognition_config(self):
|
||||||
"""Create the NVIDIA Riva ASR recognition configuration."""
|
"""Create the NVIDIA Riva ASR recognition configuration."""
|
||||||
# Create base configuration
|
# Create base configuration
|
||||||
|
s = self._settings
|
||||||
config = riva.client.RecognitionConfig(
|
config = riva.client.RecognitionConfig(
|
||||||
language_code=self._get_language_code(),
|
language_code=self._get_language_code(),
|
||||||
max_alternatives=1,
|
max_alternatives=1,
|
||||||
profanity_filter=self._settings.profanity_filter,
|
profanity_filter=s.profanity_filter,
|
||||||
enable_automatic_punctuation=self._settings.automatic_punctuation,
|
enable_automatic_punctuation=s.automatic_punctuation,
|
||||||
verbatim_transcripts=self._settings.verbatim_transcripts,
|
verbatim_transcripts=s.verbatim_transcripts,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Add word boosting if specified
|
# Add word boosting if specified
|
||||||
if self._settings.boosted_lm_words:
|
if s.boosted_lm_words:
|
||||||
riva.client.add_word_boosting_to_config(
|
riva.client.add_word_boosting_to_config(config, s.boosted_lm_words, s.boosted_lm_score)
|
||||||
config, self._settings.boosted_lm_words, self._settings.boosted_lm_score
|
|
||||||
)
|
|
||||||
|
|
||||||
# Add voice activity detection parameters
|
# Add voice activity detection parameters
|
||||||
riva.client.add_endpoint_parameters_to_config(
|
riva.client.add_endpoint_parameters_to_config(
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ Provides two STT services:
|
|||||||
|
|
||||||
import base64
|
import base64
|
||||||
import json
|
import json
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, Literal, Optional, Union
|
from typing import Any, AsyncGenerator, Literal, Optional, Union
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -35,7 +35,7 @@ from pipecat.frames.frames import (
|
|||||||
VADUserStoppedSpeakingFrame,
|
VADUserStoppedSpeakingFrame,
|
||||||
)
|
)
|
||||||
from pipecat.processors.frame_processor import FrameDirection
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.settings import STTSettings, _NotGiven, _warn_deprecated_param
|
||||||
from pipecat.services.stt_latency import OPENAI_REALTIME_TTFS_P99, OPENAI_TTFS_P99
|
from pipecat.services.stt_latency import OPENAI_REALTIME_TTFS_P99, OPENAI_TTFS_P99
|
||||||
from pipecat.services.stt_service import WebsocketSTTService
|
from pipecat.services.stt_service import WebsocketSTTService
|
||||||
from pipecat.services.whisper.base_stt import (
|
from pipecat.services.whisper.base_stt import (
|
||||||
@@ -55,6 +55,13 @@ except ModuleNotFoundError:
|
|||||||
State = None
|
State = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class OpenAISTTSettings(BaseWhisperSTTSettings):
|
||||||
|
"""Settings for the OpenAI STT service."""
|
||||||
|
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class OpenAISTTService(BaseWhisperSTTService):
|
class OpenAISTTService(BaseWhisperSTTService):
|
||||||
"""OpenAI Speech-to-Text service that generates text from audio.
|
"""OpenAI Speech-to-Text service that generates text from audio.
|
||||||
|
|
||||||
@@ -62,6 +69,8 @@ class OpenAISTTService(BaseWhisperSTTService):
|
|||||||
set via the api_key parameter or OPENAI_API_KEY environment variable.
|
set via the api_key parameter or OPENAI_API_KEY environment variable.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
_settings: OpenAISTTSettings
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -71,7 +80,7 @@ class OpenAISTTService(BaseWhisperSTTService):
|
|||||||
language: Optional[Language] = Language.EN,
|
language: Optional[Language] = Language.EN,
|
||||||
prompt: Optional[str] = None,
|
prompt: Optional[str] = None,
|
||||||
temperature: Optional[float] = None,
|
temperature: Optional[float] = None,
|
||||||
settings: Optional[BaseWhisperSTTSettings] = None,
|
settings: Optional[OpenAISTTSettings] = None,
|
||||||
ttfs_p99_latency: Optional[float] = OPENAI_TTFS_P99,
|
ttfs_p99_latency: Optional[float] = OPENAI_TTFS_P99,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
@@ -81,13 +90,25 @@ class OpenAISTTService(BaseWhisperSTTService):
|
|||||||
model: Model to use — either gpt-4o or Whisper.
|
model: Model to use — either gpt-4o or Whisper.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(model=...)`` instead.
|
Use ``settings=OpenAISTTSettings(model=...)`` instead.
|
||||||
|
|
||||||
api_key: OpenAI API key. Defaults to None.
|
api_key: OpenAI API key. Defaults to None.
|
||||||
base_url: API base URL. Defaults to None.
|
base_url: API base URL. Defaults to None.
|
||||||
language: Language of the audio input. Defaults to English.
|
language: Language of the audio input. Defaults to English.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=OpenAISTTSettings(language=...)`` instead.
|
||||||
|
|
||||||
prompt: Optional text to guide the model's style or continue a previous segment.
|
prompt: Optional text to guide the model's style or continue a previous segment.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=OpenAISTTSettings(prompt=...)`` instead.
|
||||||
|
|
||||||
temperature: Optional sampling temperature between 0 and 1. Defaults to 0.0.
|
temperature: Optional sampling temperature between 0 and 1. Defaults to 0.0.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=OpenAISTTSettings(temperature=...)`` instead.
|
||||||
|
|
||||||
settings: Runtime-updatable settings. When provided alongside deprecated
|
settings: Runtime-updatable settings. When provided alongside deprecated
|
||||||
parameters, ``settings`` values take precedence.
|
parameters, ``settings`` values take precedence.
|
||||||
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
ttfs_p99_latency: P99 latency from speech end to final transcript in seconds.
|
||||||
@@ -96,18 +117,23 @@ class OpenAISTTService(BaseWhisperSTTService):
|
|||||||
"""
|
"""
|
||||||
# --- 1. Hardcoded defaults ---
|
# --- 1. Hardcoded defaults ---
|
||||||
_language = language or Language.EN
|
_language = language or Language.EN
|
||||||
default_settings = BaseWhisperSTTSettings(
|
default_settings = OpenAISTTSettings(
|
||||||
model="gpt-4o-transcribe",
|
model="gpt-4o-transcribe",
|
||||||
language=self.language_to_service_language(_language),
|
language=self.language_to_service_language(_language),
|
||||||
base_url=base_url,
|
prompt=None,
|
||||||
prompt=prompt,
|
temperature=None,
|
||||||
temperature=temperature,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# --- 2. Deprecated direct-arg overrides ---
|
# --- 2. Deprecated direct-arg overrides ---
|
||||||
if model is not None:
|
if model is not None:
|
||||||
_warn_deprecated_param("model", BaseWhisperSTTSettings, "model")
|
_warn_deprecated_param("model", OpenAISTTSettings, "model")
|
||||||
default_settings.model = model
|
default_settings.model = model
|
||||||
|
if prompt is not None:
|
||||||
|
_warn_deprecated_param("prompt", OpenAISTTSettings, "prompt")
|
||||||
|
default_settings.prompt = prompt
|
||||||
|
if temperature is not None:
|
||||||
|
_warn_deprecated_param("temperature", OpenAISTTSettings, "temperature")
|
||||||
|
default_settings.temperature = temperature
|
||||||
|
|
||||||
# --- 3. (no params object for this service) ---
|
# --- 3. (no params object for this service) ---
|
||||||
|
|
||||||
@@ -124,7 +150,7 @@ class OpenAISTTService(BaseWhisperSTTService):
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def _transcribe(self, audio: bytes) -> Transcription:
|
async def _transcribe(self, audio: bytes) -> Transcription:
|
||||||
assert self._language is not None # Assigned in the BaseWhisperSTTService class
|
assert self._settings.language is not None
|
||||||
|
|
||||||
# Build kwargs dict with only set parameters
|
# Build kwargs dict with only set parameters
|
||||||
kwargs = {
|
kwargs = {
|
||||||
@@ -162,7 +188,7 @@ class OpenAIRealtimeSTTSettings(STTSettings):
|
|||||||
prompt: Optional prompt text to guide transcription style.
|
prompt: Optional prompt text to guide transcription style.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
prompt: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
prompt: str | None | _NotGiven = None
|
||||||
|
|
||||||
|
|
||||||
class OpenAIRealtimeSTTService(WebsocketSTTService):
|
class OpenAIRealtimeSTTService(WebsocketSTTService):
|
||||||
@@ -228,8 +254,16 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
|||||||
base_url: WebSocket base URL for the Realtime API.
|
base_url: WebSocket base URL for the Realtime API.
|
||||||
Defaults to ``"wss://api.openai.com/v1/realtime"``.
|
Defaults to ``"wss://api.openai.com/v1/realtime"``.
|
||||||
language: Language of the audio input. Defaults to English.
|
language: Language of the audio input. Defaults to English.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=OpenAIRealtimeSTTSettings(language=...)`` instead.
|
||||||
|
|
||||||
prompt: Optional prompt text to guide transcription style
|
prompt: Optional prompt text to guide transcription style
|
||||||
or provide keyword hints.
|
or provide keyword hints.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.105
|
||||||
|
Use ``settings=OpenAIRealtimeSTTSettings(prompt=...)`` instead.
|
||||||
|
|
||||||
turn_detection: Server-side VAD configuration. Defaults to
|
turn_detection: Server-side VAD configuration. Defaults to
|
||||||
``False`` (disabled), which relies on a local VAD
|
``False`` (disabled), which relies on a local VAD
|
||||||
processor in the pipeline. Pass ``None`` to use server
|
processor in the pipeline. Pass ``None`` to use server
|
||||||
@@ -257,14 +291,20 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
|||||||
# --- 1. Hardcoded defaults ---
|
# --- 1. Hardcoded defaults ---
|
||||||
default_settings = OpenAIRealtimeSTTSettings(
|
default_settings = OpenAIRealtimeSTTSettings(
|
||||||
model="gpt-4o-transcribe",
|
model="gpt-4o-transcribe",
|
||||||
language=language,
|
language=Language.EN,
|
||||||
prompt=prompt,
|
prompt=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# --- 2. Deprecated direct-arg overrides ---
|
# --- 2. Deprecated direct-arg overrides ---
|
||||||
if model is not None:
|
if model is not None:
|
||||||
_warn_deprecated_param("model", OpenAIRealtimeSTTSettings, "model")
|
_warn_deprecated_param("model", OpenAIRealtimeSTTSettings, "model")
|
||||||
default_settings.model = model
|
default_settings.model = model
|
||||||
|
if language is not None and language != Language.EN:
|
||||||
|
_warn_deprecated_param("language", OpenAIRealtimeSTTSettings, "language")
|
||||||
|
default_settings.language = language
|
||||||
|
if prompt is not None:
|
||||||
|
_warn_deprecated_param("prompt", OpenAIRealtimeSTTSettings, "prompt")
|
||||||
|
default_settings.prompt = prompt
|
||||||
|
|
||||||
# --- 3. (no params object for this service) ---
|
# --- 3. (no params object for this service) ---
|
||||||
|
|
||||||
@@ -281,7 +321,6 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
|||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._base_url = base_url
|
self._base_url = base_url
|
||||||
|
|
||||||
self._prompt = self._settings.prompt
|
|
||||||
self._turn_detection = turn_detection
|
self._turn_detection = turn_detection
|
||||||
self._noise_reduction = noise_reduction
|
self._noise_reduction = noise_reduction
|
||||||
self._should_interrupt = should_interrupt
|
self._should_interrupt = should_interrupt
|
||||||
@@ -318,8 +357,7 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
|||||||
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
||||||
"""Apply a settings delta and send session update if needed.
|
"""Apply a settings delta and send session update if needed.
|
||||||
|
|
||||||
Keeps ``_language_code`` and ``_prompt`` in sync with settings
|
Sends a ``session.update`` to the server when the session is active.
|
||||||
and sends a ``session.update`` to the server when the session is active.
|
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
delta: A :class:`STTSettings` (or ``OpenAIRealtimeSTTSettings``) delta.
|
delta: A :class:`STTSettings` (or ``OpenAIRealtimeSTTSettings``) delta.
|
||||||
@@ -329,13 +367,7 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
|||||||
"""
|
"""
|
||||||
changed = await super()._update_settings(delta)
|
changed = await super()._update_settings(delta)
|
||||||
|
|
||||||
if not changed:
|
if changed and self._session_ready:
|
||||||
return changed
|
|
||||||
|
|
||||||
if "prompt" in changed and isinstance(self._settings, OpenAIRealtimeSTTSettings):
|
|
||||||
self._prompt = self._settings.prompt
|
|
||||||
|
|
||||||
if self._session_ready:
|
|
||||||
await self._send_session_update()
|
await self._send_session_update()
|
||||||
|
|
||||||
return changed
|
return changed
|
||||||
@@ -492,8 +524,8 @@ class OpenAIRealtimeSTTService(WebsocketSTTService):
|
|||||||
if language_code:
|
if language_code:
|
||||||
transcription["language"] = language_code
|
transcription["language"] = language_code
|
||||||
|
|
||||||
if self._prompt:
|
if self._settings.prompt:
|
||||||
transcription["prompt"] = self._prompt
|
transcription["prompt"] = self._settings.prompt
|
||||||
|
|
||||||
input_audio: dict = {
|
input_audio: dict = {
|
||||||
"format": {
|
"format": {
|
||||||
|
|||||||
@@ -161,6 +161,8 @@ class OpenAITTSService(TTSService):
|
|||||||
model="gpt-4o-mini-tts",
|
model="gpt-4o-mini-tts",
|
||||||
voice="alloy",
|
voice="alloy",
|
||||||
language=None,
|
language=None,
|
||||||
|
instructions=None,
|
||||||
|
speed=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Apply direct init arg overrides (deprecated)
|
# 2. Apply direct init arg overrides (deprecated)
|
||||||
|
|||||||
@@ -6,6 +6,7 @@
|
|||||||
|
|
||||||
"""SambaNova's Speech-to-Text service implementation for real-time transcription."""
|
"""SambaNova's Speech-to-Text service implementation for real-time transcription."""
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -20,6 +21,13 @@ from pipecat.services.whisper.base_stt import (
|
|||||||
from pipecat.transcriptions.language import Language
|
from pipecat.transcriptions.language import Language
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SambaNovaSTTSettings(BaseWhisperSTTSettings):
|
||||||
|
"""Settings for the SambaNova STT service."""
|
||||||
|
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class SambaNovaSTTService(BaseWhisperSTTService): # type: ignore
|
class SambaNovaSTTService(BaseWhisperSTTService): # type: ignore
|
||||||
"""SambaNova Whisper speech-to-text service.
|
"""SambaNova Whisper speech-to-text service.
|
||||||
|
|
||||||
@@ -36,7 +44,7 @@ class SambaNovaSTTService(BaseWhisperSTTService): # type: ignore
|
|||||||
language: Optional[Language] = None,
|
language: Optional[Language] = None,
|
||||||
prompt: Optional[str] = None,
|
prompt: Optional[str] = None,
|
||||||
temperature: Optional[float] = None,
|
temperature: Optional[float] = None,
|
||||||
settings: Optional[BaseWhisperSTTSettings] = None,
|
settings: Optional[SambaNovaSTTSettings] = None,
|
||||||
ttfs_p99_latency: Optional[float] = SAMBANOVA_TTFS_P99,
|
ttfs_p99_latency: Optional[float] = SAMBANOVA_TTFS_P99,
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> None:
|
) -> None:
|
||||||
@@ -46,24 +54,24 @@ class SambaNovaSTTService(BaseWhisperSTTService): # type: ignore
|
|||||||
model: Whisper model to use.
|
model: Whisper model to use.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(model=...)`` instead.
|
Use ``settings=SambaNovaSTTSettings(model=...)`` instead.
|
||||||
|
|
||||||
api_key: SambaNova API key. Defaults to None.
|
api_key: SambaNova API key. Defaults to None.
|
||||||
base_url: API base URL. Defaults to "https://api.sambanova.ai/v1".
|
base_url: API base URL. Defaults to "https://api.sambanova.ai/v1".
|
||||||
language: Language of the audio input.
|
language: Language of the audio input.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(language=...)`` instead.
|
Use ``settings=SambaNovaSTTSettings(language=...)`` instead.
|
||||||
|
|
||||||
prompt: Optional text to guide the model's style or continue a previous segment.
|
prompt: Optional text to guide the model's style or continue a previous segment.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(prompt=...)`` instead.
|
Use ``settings=SambaNovaSTTSettings(prompt=...)`` instead.
|
||||||
|
|
||||||
temperature: Optional sampling temperature between 0 and 1.
|
temperature: Optional sampling temperature between 0 and 1.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=BaseWhisperSTTSettings(temperature=...)`` instead.
|
Use ``settings=SambaNovaSTTSettings(temperature=...)`` instead.
|
||||||
|
|
||||||
settings: Runtime-updatable settings. When provided alongside deprecated
|
settings: Runtime-updatable settings. When provided alongside deprecated
|
||||||
parameters, ``settings`` values take precedence.
|
parameters, ``settings`` values take precedence.
|
||||||
@@ -72,24 +80,25 @@ class SambaNovaSTTService(BaseWhisperSTTService): # type: ignore
|
|||||||
**kwargs: Additional arguments passed to `pipecat.services.whisper.base_stt.BaseWhisperSTTService`.
|
**kwargs: Additional arguments passed to `pipecat.services.whisper.base_stt.BaseWhisperSTTService`.
|
||||||
"""
|
"""
|
||||||
# --- 1. Hardcoded defaults ---
|
# --- 1. Hardcoded defaults ---
|
||||||
default_settings = BaseWhisperSTTSettings(
|
default_settings = SambaNovaSTTSettings(
|
||||||
model="Whisper-Large-v3",
|
model="Whisper-Large-v3",
|
||||||
language=self.language_to_service_language(Language.EN),
|
language=self.language_to_service_language(Language.EN),
|
||||||
base_url=base_url,
|
prompt=None,
|
||||||
|
temperature=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# --- 2. Deprecated direct-arg overrides ---
|
# --- 2. Deprecated direct-arg overrides ---
|
||||||
if model is not None:
|
if model is not None:
|
||||||
_warn_deprecated_param("model", BaseWhisperSTTSettings, "model")
|
_warn_deprecated_param("model", SambaNovaSTTSettings, "model")
|
||||||
default_settings.model = model
|
default_settings.model = model
|
||||||
if language is not None:
|
if language is not None:
|
||||||
_warn_deprecated_param("language", BaseWhisperSTTSettings, "language")
|
_warn_deprecated_param("language", SambaNovaSTTSettings, "language")
|
||||||
default_settings.language = self.language_to_service_language(language)
|
default_settings.language = self.language_to_service_language(language)
|
||||||
if prompt is not None:
|
if prompt is not None:
|
||||||
_warn_deprecated_param("prompt", BaseWhisperSTTSettings, "prompt")
|
_warn_deprecated_param("prompt", SambaNovaSTTSettings, "prompt")
|
||||||
default_settings.prompt = prompt
|
default_settings.prompt = prompt
|
||||||
if temperature is not None:
|
if temperature is not None:
|
||||||
_warn_deprecated_param("temperature", BaseWhisperSTTSettings, "temperature")
|
_warn_deprecated_param("temperature", SambaNovaSTTSettings, "temperature")
|
||||||
default_settings.temperature = temperature
|
default_settings.temperature = temperature
|
||||||
|
|
||||||
# --- 3. (no params object for this service) ---
|
# --- 3. (no params object for this service) ---
|
||||||
@@ -107,7 +116,7 @@ class SambaNovaSTTService(BaseWhisperSTTService): # type: ignore
|
|||||||
)
|
)
|
||||||
|
|
||||||
async def _transcribe(self, audio: bytes) -> Transcription:
|
async def _transcribe(self, audio: bytes) -> Transcription:
|
||||||
assert self._language is not None # Assigned in the BaseWhisperSTTService class
|
assert self._settings.language is not None
|
||||||
|
|
||||||
if self._include_prob_metrics:
|
if self._include_prob_metrics:
|
||||||
# https://docs.sambanova.ai/docs/en/features/audio#request-parameters
|
# https://docs.sambanova.ai/docs/en/features/audio#request-parameters
|
||||||
@@ -122,13 +131,13 @@ class SambaNovaSTTService(BaseWhisperSTTService): # type: ignore
|
|||||||
"file": ("audio.wav", audio, "audio/wav"),
|
"file": ("audio.wav", audio, "audio/wav"),
|
||||||
"model": self._settings.model,
|
"model": self._settings.model,
|
||||||
"response_format": "json",
|
"response_format": "json",
|
||||||
"language": self._language,
|
"language": self._settings.language,
|
||||||
}
|
}
|
||||||
|
|
||||||
if self._prompt is not None:
|
if self._settings.prompt is not None:
|
||||||
kwargs["prompt"] = self._prompt
|
kwargs["prompt"] = self._settings.prompt
|
||||||
|
|
||||||
if self._temperature is not None:
|
if self._settings.temperature is not None:
|
||||||
kwargs["temperature"] = self._temperature
|
kwargs["temperature"] = self._settings.temperature
|
||||||
|
|
||||||
return await self._client.audio.transcriptions.create(**kwargs)
|
return await self._client.audio.transcriptions.create(**kwargs)
|
||||||
|
|||||||
@@ -142,14 +142,13 @@ class SarvamSTTSettings(STTSettings):
|
|||||||
"""Settings for the Sarvam STT service.
|
"""Settings for the Sarvam STT service.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
prompt: Optional prompt to guide transcription/translation style.
|
prompt: Optional prompt to guide transcription/translation style/context.
|
||||||
mode: Mode of operation (transcribe, translate, verbatim, etc.).
|
Only applicable to models that support prompts (e.g., saaras:v2.5).
|
||||||
vad_signals: Enable VAD signals in response.
|
vad_signals: Enable VAD signals in response.
|
||||||
high_vad_sensitivity: Enable high VAD sensitivity.
|
high_vad_sensitivity: Enable high VAD sensitivity.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
prompt: str | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
prompt: str | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
mode: str | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
vad_signals: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
vad_signals: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
high_vad_sensitivity: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
high_vad_sensitivity: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
|
||||||
@@ -204,6 +203,9 @@ class SarvamSTTService(STTService):
|
|||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
model: Optional[str] = None,
|
model: Optional[str] = None,
|
||||||
|
mode: Optional[
|
||||||
|
Literal["transcribe", "translate", "verbatim", "translit", "codemix"]
|
||||||
|
] = None,
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
input_audio_codec: str = "wav",
|
input_audio_codec: str = "wav",
|
||||||
params: Optional[InputParams] = None,
|
params: Optional[InputParams] = None,
|
||||||
@@ -222,6 +224,9 @@ class SarvamSTTService(STTService):
|
|||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=SarvamSTTSettings(model=...)`` instead.
|
Use ``settings=SarvamSTTSettings(model=...)`` instead.
|
||||||
|
|
||||||
|
mode: Mode of operation. Options: transcribe, translate, verbatim,
|
||||||
|
translit, codemix. Only applicable to models that support it
|
||||||
|
(e.g., saaras:v3). Defaults to the model's default mode.
|
||||||
sample_rate: Audio sample rate. Defaults to 16000 if not specified.
|
sample_rate: Audio sample rate. Defaults to 16000 if not specified.
|
||||||
input_audio_codec: Audio codec/format of the input file. Defaults to "wav".
|
input_audio_codec: Audio codec/format of the input file. Defaults to "wav".
|
||||||
params: Configuration parameters for Sarvam STT service.
|
params: Configuration parameters for Sarvam STT service.
|
||||||
@@ -238,32 +243,32 @@ class SarvamSTTService(STTService):
|
|||||||
keepalive_interval: Seconds between idle checks when keepalive is enabled.
|
keepalive_interval: Seconds between idle checks when keepalive is enabled.
|
||||||
**kwargs: Additional arguments passed to the parent STTService.
|
**kwargs: Additional arguments passed to the parent STTService.
|
||||||
"""
|
"""
|
||||||
# 1. Initialize default_settings with hardcoded defaults
|
# --- 1. Hardcoded defaults ---
|
||||||
default_settings = SarvamSTTSettings(
|
default_settings = SarvamSTTSettings(
|
||||||
model="saarika:v2.5",
|
model="saarika:v2.5",
|
||||||
language=None,
|
language=None,
|
||||||
prompt=None,
|
prompt=None,
|
||||||
mode=None,
|
|
||||||
vad_signals=None,
|
vad_signals=None,
|
||||||
high_vad_sensitivity=None,
|
high_vad_sensitivity=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Apply direct init arg overrides (deprecated)
|
# --- 2. Deprecated direct-arg overrides ---
|
||||||
if model is not None:
|
if model is not None:
|
||||||
_warn_deprecated_param("model", SarvamSTTSettings, "model")
|
_warn_deprecated_param("model", SarvamSTTSettings, "model")
|
||||||
default_settings.model = model
|
default_settings.model = model
|
||||||
|
|
||||||
# 3. Apply params overrides — only if settings not provided
|
# --- 3. Deprecated params overrides ---
|
||||||
if params is not None:
|
if params is not None:
|
||||||
_warn_deprecated_param("params", SarvamSTTSettings)
|
_warn_deprecated_param("params", SarvamSTTSettings)
|
||||||
if not settings:
|
if not settings:
|
||||||
default_settings.language = params.language
|
default_settings.language = params.language
|
||||||
default_settings.prompt = params.prompt
|
default_settings.prompt = params.prompt
|
||||||
default_settings.mode = params.mode
|
if params.mode is not None:
|
||||||
|
mode = params.mode
|
||||||
default_settings.vad_signals = params.vad_signals
|
default_settings.vad_signals = params.vad_signals
|
||||||
default_settings.high_vad_sensitivity = params.high_vad_sensitivity
|
default_settings.high_vad_sensitivity = params.high_vad_sensitivity
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# --- 4. Settings delta (canonical API, always wins) ---
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
default_settings.apply_update(settings)
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
@@ -278,7 +283,7 @@ class SarvamSTTService(STTService):
|
|||||||
# Validate parameters against model capabilities
|
# Validate parameters against model capabilities
|
||||||
if default_settings.prompt is not None and not self._config.supports_prompt:
|
if default_settings.prompt is not None and not self._config.supports_prompt:
|
||||||
raise ValueError(f"Model '{resolved_model}' does not support prompt parameter.")
|
raise ValueError(f"Model '{resolved_model}' does not support prompt parameter.")
|
||||||
if default_settings.mode is not None and not self._config.supports_mode:
|
if mode is not None and not self._config.supports_mode:
|
||||||
raise ValueError(f"Model '{resolved_model}' does not support mode parameter.")
|
raise ValueError(f"Model '{resolved_model}' does not support mode parameter.")
|
||||||
if default_settings.language is not None and not self._config.supports_language:
|
if default_settings.language is not None and not self._config.supports_language:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -286,8 +291,8 @@ class SarvamSTTService(STTService):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Resolve mode default from model config
|
# Resolve mode default from model config
|
||||||
if default_settings.mode is None:
|
if mode is None:
|
||||||
default_settings.mode = self._config.default_mode
|
mode = self._config.default_mode
|
||||||
|
|
||||||
super().__init__(
|
super().__init__(
|
||||||
sample_rate=sample_rate,
|
sample_rate=sample_rate,
|
||||||
@@ -300,6 +305,9 @@ class SarvamSTTService(STTService):
|
|||||||
|
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
|
|
||||||
|
# Init-only connection config (not runtime-updatable)
|
||||||
|
self._mode = mode
|
||||||
|
|
||||||
# Store connection parameters
|
# Store connection parameters
|
||||||
self._input_audio_codec = input_audio_codec
|
self._input_audio_codec = input_audio_codec
|
||||||
|
|
||||||
@@ -380,30 +388,26 @@ class SarvamSTTService(STTService):
|
|||||||
f"Model '{self._settings.model}' does not support language parameter "
|
f"Model '{self._settings.model}' does not support language parameter "
|
||||||
"(auto-detects language)."
|
"(auto-detects language)."
|
||||||
)
|
)
|
||||||
|
if (
|
||||||
if isinstance(delta, SarvamSTTSettings):
|
isinstance(delta, SarvamSTTSettings)
|
||||||
if is_given(delta.prompt) and delta.prompt is not None:
|
and is_given(delta.prompt)
|
||||||
if not self._config.supports_prompt:
|
and delta.prompt is not None
|
||||||
raise ValueError(
|
):
|
||||||
f"Model '{self._settings.model}' does not support prompt parameter."
|
if not self._config.supports_prompt:
|
||||||
)
|
raise ValueError(
|
||||||
if is_given(delta.mode) and delta.mode is not None:
|
f"Model '{self._settings.model}' does not support prompt parameter."
|
||||||
if not self._config.supports_mode:
|
)
|
||||||
raise ValueError(
|
|
||||||
f"Model '{self._settings.model}' does not support mode parameter."
|
|
||||||
)
|
|
||||||
|
|
||||||
changed = await super()._update_settings(delta)
|
changed = await super()._update_settings(delta)
|
||||||
|
|
||||||
# TODO: someday we could reconnect here to apply updated settings.
|
# Prompt is a WebSocket connect-time parameter; reconnect to apply.
|
||||||
# Code might look something like the below:
|
if "prompt" in changed:
|
||||||
# if not changed:
|
await self._disconnect()
|
||||||
# return changed
|
await self._connect()
|
||||||
|
|
||||||
# await self._disconnect()
|
unhandled = {k: v for k, v in changed.items() if k != "prompt"}
|
||||||
# await self._connect()
|
if unhandled:
|
||||||
|
self._warn_unhandled_updated_settings(unhandled)
|
||||||
self._warn_unhandled_updated_settings(changed)
|
|
||||||
|
|
||||||
return changed
|
return changed
|
||||||
|
|
||||||
@@ -542,8 +546,8 @@ class SarvamSTTService(STTService):
|
|||||||
connect_kwargs["language_code"] = language_string
|
connect_kwargs["language_code"] = language_string
|
||||||
|
|
||||||
# Add mode for models that support it
|
# Add mode for models that support it
|
||||||
if self._config.supports_mode and self._settings.mode is not None:
|
if self._config.supports_mode and self._mode is not None:
|
||||||
connect_kwargs["mode"] = self._settings.mode
|
connect_kwargs["mode"] = self._mode
|
||||||
|
|
||||||
# Prompt support differs across sarvamai versions. Prefer connect-time prompt
|
# Prompt support differs across sarvamai versions. Prefer connect-time prompt
|
||||||
# when available and gracefully degrade if the SDK doesn't accept it.
|
# when available and gracefully degrade if the SDK doesn't accept it.
|
||||||
|
|||||||
@@ -144,8 +144,6 @@ class SonioxSTTSettings(STTSettings):
|
|||||||
"""Settings for Soniox STT service.
|
"""Settings for Soniox STT service.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
audio_format: Audio format to use for transcription.
|
|
||||||
num_channels: Number of channels to use for transcription.
|
|
||||||
language_hints: List of language hints to use for transcription.
|
language_hints: List of language hints to use for transcription.
|
||||||
language_hints_strict: If true, strictly enforce language hints.
|
language_hints_strict: If true, strictly enforce language hints.
|
||||||
context: Customization for transcription. String for models with
|
context: Customization for transcription. String for models with
|
||||||
@@ -156,8 +154,6 @@ class SonioxSTTSettings(STTSettings):
|
|||||||
client_reference_id: Client reference ID to use for transcription.
|
client_reference_id: Client reference ID to use for transcription.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
audio_format: str | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
num_channels: int | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
language_hints: List[Language] | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
language_hints: List[Language] | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
language_hints_strict: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
language_hints_strict: bool | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
context: SonioxContextObject | str | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
context: SonioxContextObject | str | None | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
@@ -187,6 +183,8 @@ class SonioxSTTService(WebsocketSTTService):
|
|||||||
url: str = "wss://stt-rt.soniox.com/transcribe-websocket",
|
url: str = "wss://stt-rt.soniox.com/transcribe-websocket",
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
model: Optional[str] = None,
|
model: Optional[str] = None,
|
||||||
|
audio_format: str = "pcm_s16le",
|
||||||
|
num_channels: int = 1,
|
||||||
params: Optional[SonioxInputParams] = None,
|
params: Optional[SonioxInputParams] = None,
|
||||||
vad_force_turn_endpoint: bool = True,
|
vad_force_turn_endpoint: bool = True,
|
||||||
settings: Optional[SonioxSTTSettings] = None,
|
settings: Optional[SonioxSTTSettings] = None,
|
||||||
@@ -204,6 +202,8 @@ class SonioxSTTService(WebsocketSTTService):
|
|||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
Use ``settings=SonioxSTTSettings(model=...)`` instead.
|
Use ``settings=SonioxSTTSettings(model=...)`` instead.
|
||||||
|
|
||||||
|
audio_format: Audio format for transcription. Defaults to ``"pcm_s16le"``.
|
||||||
|
num_channels: Number of audio channels. Defaults to 1.
|
||||||
params: Additional configuration parameters, such as language hints, context and
|
params: Additional configuration parameters, such as language hints, context and
|
||||||
speaker diarization.
|
speaker diarization.
|
||||||
|
|
||||||
@@ -218,12 +218,10 @@ class SonioxSTTService(WebsocketSTTService):
|
|||||||
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
Override for your deployment. See https://github.com/pipecat-ai/stt-benchmark
|
||||||
**kwargs: Additional arguments passed to the STTService.
|
**kwargs: Additional arguments passed to the STTService.
|
||||||
"""
|
"""
|
||||||
# 1. Initialize default_settings with hardcoded defaults
|
# --- 1. Hardcoded defaults ---
|
||||||
default_settings = SonioxSTTSettings(
|
default_settings = SonioxSTTSettings(
|
||||||
model="stt-rt-v4",
|
model="stt-rt-v4",
|
||||||
language=None,
|
language=None,
|
||||||
audio_format="pcm_s16le",
|
|
||||||
num_channels=1,
|
|
||||||
language_hints=None,
|
language_hints=None,
|
||||||
language_hints_strict=None,
|
language_hints_strict=None,
|
||||||
context=None,
|
context=None,
|
||||||
@@ -232,18 +230,20 @@ class SonioxSTTService(WebsocketSTTService):
|
|||||||
client_reference_id=None,
|
client_reference_id=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. Apply direct init arg overrides (deprecated)
|
# --- 2. Deprecated direct-arg overrides ---
|
||||||
if model is not None:
|
if model is not None:
|
||||||
_warn_deprecated_param("model", SonioxSTTSettings, "model")
|
_warn_deprecated_param("model", SonioxSTTSettings, "model")
|
||||||
default_settings.model = model
|
default_settings.model = model
|
||||||
|
|
||||||
# 3. Apply params overrides — only if settings not provided
|
# --- 3. Deprecated params overrides ---
|
||||||
if params is not None:
|
if params is not None:
|
||||||
_warn_deprecated_param("params", SonioxSTTSettings)
|
_warn_deprecated_param("params", SonioxSTTSettings)
|
||||||
if not settings:
|
if not settings:
|
||||||
default_settings.model = params.model
|
default_settings.model = params.model
|
||||||
default_settings.audio_format = params.audio_format
|
if params.audio_format is not None:
|
||||||
default_settings.num_channels = params.num_channels
|
audio_format = params.audio_format
|
||||||
|
if params.num_channels is not None:
|
||||||
|
num_channels = params.num_channels
|
||||||
default_settings.language_hints = params.language_hints
|
default_settings.language_hints = params.language_hints
|
||||||
default_settings.language_hints_strict = params.language_hints_strict
|
default_settings.language_hints_strict = params.language_hints_strict
|
||||||
default_settings.context = params.context
|
default_settings.context = params.context
|
||||||
@@ -253,7 +253,7 @@ class SonioxSTTService(WebsocketSTTService):
|
|||||||
)
|
)
|
||||||
default_settings.client_reference_id = params.client_reference_id
|
default_settings.client_reference_id = params.client_reference_id
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# --- 4. Settings delta (canonical API, always wins) ---
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
default_settings.apply_update(settings)
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
@@ -270,6 +270,10 @@ class SonioxSTTService(WebsocketSTTService):
|
|||||||
self._url = url
|
self._url = url
|
||||||
self._vad_force_turn_endpoint = vad_force_turn_endpoint
|
self._vad_force_turn_endpoint = vad_force_turn_endpoint
|
||||||
|
|
||||||
|
# Init-only audio config
|
||||||
|
self._audio_format = audio_format
|
||||||
|
self._num_channels = num_channels
|
||||||
|
|
||||||
self._final_transcription_buffer = []
|
self._final_transcription_buffer = []
|
||||||
self._last_tokens_received: Optional[float] = None
|
self._last_tokens_received: Optional[float] = None
|
||||||
|
|
||||||
@@ -438,8 +442,8 @@ class SonioxSTTService(WebsocketSTTService):
|
|||||||
config = {
|
config = {
|
||||||
"api_key": self._api_key,
|
"api_key": self._api_key,
|
||||||
"model": s.model,
|
"model": s.model,
|
||||||
"audio_format": s.audio_format,
|
"audio_format": self._audio_format,
|
||||||
"num_channels": s.num_channels or 1,
|
"num_channels": self._num_channels,
|
||||||
"enable_endpoint_detection": enable_endpoint_detection,
|
"enable_endpoint_detection": enable_endpoint_detection,
|
||||||
"sample_rate": self.sample_rate,
|
"sample_rate": self.sample_rate,
|
||||||
"language_hints": _prepare_language_hints(s.language_hints),
|
"language_hints": _prepare_language_hints(s.language_hints),
|
||||||
|
|||||||
@@ -100,7 +100,6 @@ class SpeechmaticsSTTSettings(STTSettings):
|
|||||||
focus_mode: Speaker focus mode for diarization.
|
focus_mode: Speaker focus mode for diarization.
|
||||||
known_speakers: List of known speaker labels and identifiers.
|
known_speakers: List of known speaker labels and identifiers.
|
||||||
additional_vocab: List of additional vocabulary entries.
|
additional_vocab: List of additional vocabulary entries.
|
||||||
audio_encoding: Audio encoding format.
|
|
||||||
operating_point: Operating point for accuracy vs. latency.
|
operating_point: Operating point for accuracy vs. latency.
|
||||||
max_delay: Maximum delay in seconds for transcription.
|
max_delay: Maximum delay in seconds for transcription.
|
||||||
end_of_utterance_silence_trigger: Maximum delay for end of utterance trigger.
|
end_of_utterance_silence_trigger: Maximum delay for end of utterance trigger.
|
||||||
@@ -126,7 +125,6 @@ class SpeechmaticsSTTSettings(STTSettings):
|
|||||||
additional_vocab: list[AdditionalVocabEntry] | _NotGiven = field(
|
additional_vocab: list[AdditionalVocabEntry] | _NotGiven = field(
|
||||||
default_factory=lambda: NOT_GIVEN
|
default_factory=lambda: NOT_GIVEN
|
||||||
)
|
)
|
||||||
audio_encoding: AudioEncoding | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
operating_point: OperatingPoint | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
operating_point: OperatingPoint | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
max_delay: float | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
max_delay: float | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
end_of_utterance_silence_trigger: float | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
end_of_utterance_silence_trigger: float | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
@@ -344,8 +342,8 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
class UpdateParams(BaseModel):
|
class UpdateParams(BaseModel):
|
||||||
"""Update parameters for Speechmatics STT service.
|
"""Update parameters for Speechmatics STT service.
|
||||||
|
|
||||||
These are the only parameters that can be changed once a session has started. If you need to
|
.. deprecated:: 0.0.104
|
||||||
change the language, etc., then you must create a new instance of the service.
|
Use ``SpeechmaticsSTTSettings`` with ``STTUpdateSettingsFrame`` instead.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
focus_speakers: List of speaker IDs to focus on. When enabled, only these speakers are
|
focus_speakers: List of speaker IDs to focus on. When enabled, only these speakers are
|
||||||
@@ -379,6 +377,7 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
api_key: str | None = None,
|
api_key: str | None = None,
|
||||||
base_url: str | None = None,
|
base_url: str | None = None,
|
||||||
sample_rate: int | None = None,
|
sample_rate: int | None = None,
|
||||||
|
encoding: AudioEncoding = AudioEncoding.PCM_S16LE,
|
||||||
params: InputParams | None = None,
|
params: InputParams | None = None,
|
||||||
should_interrupt: bool = True,
|
should_interrupt: bool = True,
|
||||||
settings: SpeechmaticsSTTSettings | None = None,
|
settings: SpeechmaticsSTTSettings | None = None,
|
||||||
@@ -393,6 +392,7 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
base_url: Base URL for Speechmatics API. Uses environment variable `SPEECHMATICS_RT_URL`
|
base_url: Base URL for Speechmatics API. Uses environment variable `SPEECHMATICS_RT_URL`
|
||||||
or defaults to `wss://eu2.rt.speechmatics.com/v2`.
|
or defaults to `wss://eu2.rt.speechmatics.com/v2`.
|
||||||
sample_rate: Optional audio sample rate in Hz.
|
sample_rate: Optional audio sample rate in Hz.
|
||||||
|
encoding: Audio encoding format. Defaults to ``AudioEncoding.PCM_S16LE``.
|
||||||
params: Input parameters for the service.
|
params: Input parameters for the service.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
@@ -423,7 +423,7 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
_params = params or SpeechmaticsSTTService.InputParams()
|
_params = params or SpeechmaticsSTTService.InputParams()
|
||||||
self._check_deprecated_args(kwargs, _params)
|
self._check_deprecated_args(kwargs, _params)
|
||||||
|
|
||||||
# 1. Initialize default_settings with hardcoded defaults
|
# --- 1. Hardcoded defaults ---
|
||||||
default_settings = SpeechmaticsSTTSettings(
|
default_settings = SpeechmaticsSTTSettings(
|
||||||
model=None, # Will be resolved from operating_point after config is built
|
model=None, # Will be resolved from operating_point after config is built
|
||||||
language=Language.EN,
|
language=Language.EN,
|
||||||
@@ -436,7 +436,6 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
focus_mode=SpeakerFocusMode.RETAIN,
|
focus_mode=SpeakerFocusMode.RETAIN,
|
||||||
known_speakers=[],
|
known_speakers=[],
|
||||||
additional_vocab=[],
|
additional_vocab=[],
|
||||||
audio_encoding=AudioEncoding.PCM_S16LE,
|
|
||||||
operating_point=None,
|
operating_point=None,
|
||||||
max_delay=None,
|
max_delay=None,
|
||||||
end_of_utterance_silence_trigger=None,
|
end_of_utterance_silence_trigger=None,
|
||||||
@@ -451,9 +450,9 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
extra_params=None,
|
extra_params=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
# 2. No direct init arg overrides
|
# --- 2. No direct init arg overrides ---
|
||||||
|
|
||||||
# 3. Apply params overrides — only if settings not provided
|
# --- 3. Deprecated params overrides ---
|
||||||
if params is not None:
|
if params is not None:
|
||||||
_warn_deprecated_param("params", SpeechmaticsSTTSettings)
|
_warn_deprecated_param("params", SpeechmaticsSTTSettings)
|
||||||
if not settings:
|
if not settings:
|
||||||
@@ -475,7 +474,7 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
default_settings.focus_mode = _params.focus_mode
|
default_settings.focus_mode = _params.focus_mode
|
||||||
default_settings.known_speakers = _params.known_speakers
|
default_settings.known_speakers = _params.known_speakers
|
||||||
default_settings.additional_vocab = _params.additional_vocab
|
default_settings.additional_vocab = _params.additional_vocab
|
||||||
default_settings.audio_encoding = _params.audio_encoding
|
encoding = _params.audio_encoding
|
||||||
default_settings.operating_point = _params.operating_point
|
default_settings.operating_point = _params.operating_point
|
||||||
default_settings.max_delay = _params.max_delay
|
default_settings.max_delay = _params.max_delay
|
||||||
default_settings.end_of_utterance_silence_trigger = (
|
default_settings.end_of_utterance_silence_trigger = (
|
||||||
@@ -493,10 +492,11 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
|
|
||||||
# Build SDK config from settings, then resolve model from operating_point
|
# Build SDK config from settings, then resolve model from operating_point
|
||||||
self._client: VoiceAgentClient | None = None
|
self._client: VoiceAgentClient | None = None
|
||||||
|
self._audio_encoding = encoding
|
||||||
self._config: VoiceAgentConfig = self._build_config(default_settings)
|
self._config: VoiceAgentConfig = self._build_config(default_settings)
|
||||||
default_settings.model = self._config.operating_point.value
|
default_settings.model = self._config.operating_point.value
|
||||||
|
|
||||||
# 4. Apply settings delta (canonical API, always wins)
|
# --- 4. Settings delta (canonical API, always wins) ---
|
||||||
if settings is not None:
|
if settings is not None:
|
||||||
default_settings.apply_update(settings)
|
default_settings.apply_update(settings)
|
||||||
|
|
||||||
@@ -720,6 +720,9 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
# Preset from turn detection mode
|
# Preset from turn detection mode
|
||||||
config = VoiceAgentConfigPreset.load(s.turn_detection_mode.value)
|
config = VoiceAgentConfigPreset.load(s.turn_detection_mode.value)
|
||||||
|
|
||||||
|
# Audio encoding (init-only, stored as instance attribute)
|
||||||
|
config.audio_encoding = self._audio_encoding
|
||||||
|
|
||||||
# Language + domain
|
# Language + domain
|
||||||
language = s.language
|
language = s.language
|
||||||
config.language = self._language_to_speechmatics_language(language)
|
config.language = self._language_to_speechmatics_language(language)
|
||||||
@@ -773,7 +776,7 @@ class SpeechmaticsSTTService(STTService):
|
|||||||
) -> None:
|
) -> None:
|
||||||
"""Updates the speaker configuration.
|
"""Updates the speaker configuration.
|
||||||
|
|
||||||
.. deprecated::
|
.. deprecated:: 0.0.104
|
||||||
Use ``STTUpdateSettingsFrame`` with
|
Use ``STTUpdateSettingsFrame`` with
|
||||||
``SpeechmaticsSTTSettings(...)`` instead.
|
``SpeechmaticsSTTSettings(...)`` instead.
|
||||||
|
|
||||||
|
|||||||
@@ -10,15 +10,15 @@ This module provides common functionality for services implementing the Whisper
|
|||||||
interface, including language mapping, metrics generation, and error handling.
|
interface, including language mapping, metrics generation, and error handling.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, Optional
|
from typing import AsyncGenerator, Optional
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from openai import AsyncOpenAI
|
from openai import AsyncOpenAI
|
||||||
from openai.types.audio import Transcription
|
from openai.types.audio import Transcription
|
||||||
|
|
||||||
from pipecat.frames.frames import ErrorFrame, Frame, TranscriptionFrame
|
from pipecat.frames.frames import ErrorFrame, Frame, TranscriptionFrame
|
||||||
from pipecat.services.settings import NOT_GIVEN, STTSettings, _NotGiven, _warn_deprecated_param
|
from pipecat.services.settings import STTSettings, _NotGiven, _warn_deprecated_param
|
||||||
from pipecat.services.stt_latency import WHISPER_TTFS_P99
|
from pipecat.services.stt_latency import WHISPER_TTFS_P99
|
||||||
from pipecat.services.stt_service import SegmentedSTTService
|
from pipecat.services.stt_service import SegmentedSTTService
|
||||||
from pipecat.transcriptions.language import Language, resolve_language
|
from pipecat.transcriptions.language import Language, resolve_language
|
||||||
@@ -31,15 +31,13 @@ class BaseWhisperSTTSettings(STTSettings):
|
|||||||
"""Settings for Whisper API-based STT services.
|
"""Settings for Whisper API-based STT services.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
base_url: API base URL.
|
|
||||||
prompt: Optional text to guide the model's style or continue
|
prompt: Optional text to guide the model's style or continue
|
||||||
a previous segment.
|
a previous segment.
|
||||||
temperature: Sampling temperature between 0 and 1.
|
temperature: Sampling temperature between 0 and 1.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
base_url: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
prompt: str | None | _NotGiven = None
|
||||||
prompt: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
temperature: float | None | _NotGiven = None
|
||||||
temperature: float | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
|
|
||||||
|
|
||||||
def language_to_whisper_language(language: Language) -> Optional[str]:
|
def language_to_whisper_language(language: Language) -> Optional[str]:
|
||||||
@@ -185,7 +183,6 @@ class BaseWhisperSTTService(SegmentedSTTService):
|
|||||||
default_settings = BaseWhisperSTTSettings(
|
default_settings = BaseWhisperSTTSettings(
|
||||||
model=None,
|
model=None,
|
||||||
language=None,
|
language=None,
|
||||||
base_url=base_url,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# --- 2. Deprecated direct-arg overrides ---
|
# --- 2. Deprecated direct-arg overrides ---
|
||||||
@@ -214,32 +211,12 @@ class BaseWhisperSTTService(SegmentedSTTService):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
self._client = self._create_client(api_key, base_url)
|
self._client = self._create_client(api_key, base_url)
|
||||||
self._language = self._settings.language
|
|
||||||
self._prompt = self._settings.prompt
|
|
||||||
self._temperature = self._settings.temperature
|
|
||||||
self._include_prob_metrics = include_prob_metrics
|
self._include_prob_metrics = include_prob_metrics
|
||||||
self._push_empty_transcripts = push_empty_transcripts
|
self._push_empty_transcripts = push_empty_transcripts
|
||||||
|
|
||||||
def _create_client(self, api_key: Optional[str], base_url: Optional[str]):
|
def _create_client(self, api_key: Optional[str], base_url: Optional[str]):
|
||||||
return AsyncOpenAI(api_key=api_key, base_url=base_url)
|
return AsyncOpenAI(api_key=api_key, base_url=base_url)
|
||||||
|
|
||||||
async def _update_settings(self, delta: STTSettings) -> dict[str, Any]:
|
|
||||||
"""Apply a settings delta, syncing instance variables.
|
|
||||||
|
|
||||||
Keeps ``_language``, ``_prompt``, and ``_temperature`` in sync with
|
|
||||||
the settings fields.
|
|
||||||
"""
|
|
||||||
changed = await super()._update_settings(delta)
|
|
||||||
|
|
||||||
if "language" in changed:
|
|
||||||
self._language = self._settings.language
|
|
||||||
if "prompt" in changed:
|
|
||||||
self._prompt = self._settings.prompt
|
|
||||||
if "temperature" in changed:
|
|
||||||
self._temperature = self._settings.temperature
|
|
||||||
|
|
||||||
return changed
|
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
"""Whether this service can generate processing metrics.
|
"""Whether this service can generate processing metrics.
|
||||||
|
|
||||||
@@ -289,7 +266,7 @@ class BaseWhisperSTTService(SegmentedSTTService):
|
|||||||
logger.warning("Received empty transcription from API")
|
logger.warning("Received empty transcription from API")
|
||||||
|
|
||||||
if text or self._push_empty_transcripts:
|
if text or self._push_empty_transcripts:
|
||||||
await self._handle_transcription(text, True, self._language)
|
await self._handle_transcription(text, True, self._settings.language)
|
||||||
logger.debug(f"Transcription: [{text}]")
|
logger.debug(f"Transcription: [{text}]")
|
||||||
yield TranscriptionFrame(
|
yield TranscriptionFrame(
|
||||||
text,
|
text,
|
||||||
|
|||||||
@@ -179,13 +179,9 @@ class WhisperSTTSettings(STTSettings):
|
|||||||
"""Settings for the local Whisper (Faster Whisper) STT service.
|
"""Settings for the local Whisper (Faster Whisper) STT service.
|
||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
device: Inference device ('cpu', 'cuda', or 'auto').
|
|
||||||
compute_type: Compute type for inference ('default', 'int8', etc.).
|
|
||||||
no_speech_prob: Probability threshold for filtering non-speech segments.
|
no_speech_prob: Probability threshold for filtering non-speech segments.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
device: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
compute_type: str | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
|
||||||
no_speech_prob: float | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
no_speech_prob: float | _NotGiven = field(default_factory=lambda: NOT_GIVEN)
|
||||||
|
|
||||||
|
|
||||||
@@ -217,8 +213,8 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
model: Optional[str | Model] = None,
|
model: Optional[str | Model] = None,
|
||||||
device: Optional[str] = None,
|
device: str = "auto",
|
||||||
compute_type: Optional[str] = None,
|
compute_type: str = "default",
|
||||||
no_speech_prob: Optional[float] = None,
|
no_speech_prob: Optional[float] = None,
|
||||||
language: Optional[Language] = None,
|
language: Optional[Language] = None,
|
||||||
settings: Optional[WhisperSTTSettings] = None,
|
settings: Optional[WhisperSTTSettings] = None,
|
||||||
@@ -233,15 +229,9 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
Use ``settings=WhisperSTTSettings(model=...)`` instead.
|
Use ``settings=WhisperSTTSettings(model=...)`` instead.
|
||||||
|
|
||||||
device: The device to run inference on ('cpu', 'cuda', or 'auto').
|
device: The device to run inference on ('cpu', 'cuda', or 'auto').
|
||||||
|
Defaults to ``"auto"``.
|
||||||
.. deprecated:: 0.0.105
|
compute_type: The compute type for inference ('default', 'int8',
|
||||||
Use ``settings=WhisperSTTSettings(device=...)`` instead.
|
'int8_float16', etc.). Defaults to ``"default"``.
|
||||||
|
|
||||||
compute_type: The compute type for inference ('default', 'int8', 'int8_float16', etc.).
|
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
|
||||||
Use ``settings=WhisperSTTSettings(compute_type=...)`` instead.
|
|
||||||
|
|
||||||
no_speech_prob: Probability threshold for filtering out non-speech segments.
|
no_speech_prob: Probability threshold for filtering out non-speech segments.
|
||||||
|
|
||||||
.. deprecated:: 0.0.105
|
.. deprecated:: 0.0.105
|
||||||
@@ -260,8 +250,6 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
default_settings = WhisperSTTSettings(
|
default_settings = WhisperSTTSettings(
|
||||||
model=Model.DISTIL_MEDIUM_EN.value,
|
model=Model.DISTIL_MEDIUM_EN.value,
|
||||||
language=Language.EN,
|
language=Language.EN,
|
||||||
device="auto",
|
|
||||||
compute_type="default",
|
|
||||||
no_speech_prob=0.4,
|
no_speech_prob=0.4,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -269,12 +257,6 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
if model is not None:
|
if model is not None:
|
||||||
_warn_deprecated_param("model", WhisperSTTSettings, "model")
|
_warn_deprecated_param("model", WhisperSTTSettings, "model")
|
||||||
default_settings.model = model if isinstance(model, str) else model.value
|
default_settings.model = model if isinstance(model, str) else model.value
|
||||||
if device is not None:
|
|
||||||
_warn_deprecated_param("device", WhisperSTTSettings, "device")
|
|
||||||
default_settings.device = device
|
|
||||||
if compute_type is not None:
|
|
||||||
_warn_deprecated_param("compute_type", WhisperSTTSettings, "compute_type")
|
|
||||||
default_settings.compute_type = compute_type
|
|
||||||
if no_speech_prob is not None:
|
if no_speech_prob is not None:
|
||||||
_warn_deprecated_param("no_speech_prob", WhisperSTTSettings, "no_speech_prob")
|
_warn_deprecated_param("no_speech_prob", WhisperSTTSettings, "no_speech_prob")
|
||||||
default_settings.no_speech_prob = no_speech_prob
|
default_settings.no_speech_prob = no_speech_prob
|
||||||
@@ -292,9 +274,11 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
settings=default_settings,
|
settings=default_settings,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
self._device: str = self._settings.device
|
|
||||||
self._compute_type = self._settings.compute_type
|
# Init-only inference config
|
||||||
self._no_speech_prob = self._settings.no_speech_prob
|
self._device = device
|
||||||
|
self._compute_type = compute_type
|
||||||
|
|
||||||
self._model: Optional[WhisperModel] = None
|
self._model: Optional[WhisperModel] = None
|
||||||
|
|
||||||
self._load()
|
self._load()
|
||||||
@@ -373,7 +357,7 @@ class WhisperSTTService(SegmentedSTTService):
|
|||||||
)
|
)
|
||||||
text: str = ""
|
text: str = ""
|
||||||
for segment in segments:
|
for segment in segments:
|
||||||
if segment.no_speech_prob < self._no_speech_prob:
|
if segment.no_speech_prob < self._settings.no_speech_prob:
|
||||||
text += f"{segment.text} "
|
text += f"{segment.text} "
|
||||||
|
|
||||||
await self.stop_processing_metrics()
|
await self.stop_processing_metrics()
|
||||||
@@ -471,9 +455,6 @@ class WhisperSTTServiceMLX(WhisperSTTService):
|
|||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._no_speech_prob = self._settings.no_speech_prob
|
|
||||||
self._temperature = self._settings.temperature
|
|
||||||
|
|
||||||
# No need to call _load() as MLX Whisper loads models on demand
|
# No need to call _load() as MLX Whisper loads models on demand
|
||||||
|
|
||||||
@override
|
@override
|
||||||
@@ -514,7 +495,7 @@ class WhisperSTTServiceMLX(WhisperSTTService):
|
|||||||
mlx_whisper.transcribe,
|
mlx_whisper.transcribe,
|
||||||
audio_float,
|
audio_float,
|
||||||
path_or_hf_repo=self._settings.model,
|
path_or_hf_repo=self._settings.model,
|
||||||
temperature=self._temperature,
|
temperature=self._settings.temperature,
|
||||||
language=self._settings.language,
|
language=self._settings.language,
|
||||||
)
|
)
|
||||||
text: str = ""
|
text: str = ""
|
||||||
@@ -523,7 +504,7 @@ class WhisperSTTServiceMLX(WhisperSTTService):
|
|||||||
if segment.get("compression_ratio", None) == 0.5555555555555556:
|
if segment.get("compression_ratio", None) == 0.5555555555555556:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
if segment.get("no_speech_prob", 0.0) < self._no_speech_prob:
|
if segment.get("no_speech_prob", 0.0) < self._settings.no_speech_prob:
|
||||||
text += f"{segment.get('text', '')} "
|
text += f"{segment.get('text', '')} "
|
||||||
|
|
||||||
if len(text.strip()) == 0:
|
if len(text.strip()) == 0:
|
||||||
|
|||||||
@@ -328,8 +328,6 @@ class TestDeepgramSTTSettingsApplyUpdate:
|
|||||||
defaults = dict(
|
defaults = dict(
|
||||||
model="nova-3-general",
|
model="nova-3-general",
|
||||||
language="en",
|
language="en",
|
||||||
encoding="linear16",
|
|
||||||
channels=1,
|
|
||||||
interim_results=True,
|
interim_results=True,
|
||||||
smart_format=False,
|
smart_format=False,
|
||||||
punctuate=True,
|
punctuate=True,
|
||||||
@@ -350,8 +348,8 @@ class TestDeepgramSTTSettingsApplyUpdate:
|
|||||||
assert current.punctuate is False
|
assert current.punctuate is False
|
||||||
assert "punctuate" in changed
|
assert "punctuate" in changed
|
||||||
# Other fields are untouched
|
# Other fields are untouched
|
||||||
assert current.encoding == "linear16"
|
assert current.model == "nova-3-general"
|
||||||
assert current.channels == 1
|
assert current.language == "en"
|
||||||
|
|
||||||
def test_apply_update_model(self):
|
def test_apply_update_model(self):
|
||||||
"""model field is updated directly."""
|
"""model field is updated directly."""
|
||||||
@@ -427,8 +425,6 @@ class TestDeepgramSTTSettingsFromMapping:
|
|||||||
current = DeepgramSTTSettings(
|
current = DeepgramSTTSettings(
|
||||||
model="nova-3-general",
|
model="nova-3-general",
|
||||||
language="en",
|
language="en",
|
||||||
encoding="linear16",
|
|
||||||
channels=1,
|
|
||||||
interim_results=True,
|
interim_results=True,
|
||||||
punctuate=True,
|
punctuate=True,
|
||||||
profanity_filter=True,
|
profanity_filter=True,
|
||||||
@@ -442,7 +438,6 @@ class TestDeepgramSTTSettingsFromMapping:
|
|||||||
assert current.punctuate is False
|
assert current.punctuate is False
|
||||||
assert current.diarize is True
|
assert current.diarize is True
|
||||||
# Unchanged fields stay put
|
# Unchanged fields stay put
|
||||||
assert current.encoding == "linear16"
|
|
||||||
assert current.model == "nova-3-general"
|
assert current.model == "nova-3-general"
|
||||||
assert "punctuate" in changed
|
assert "punctuate" in changed
|
||||||
|
|
||||||
@@ -451,8 +446,6 @@ class TestDeepgramSTTSettingsFromMapping:
|
|||||||
current = DeepgramSTTSettings(
|
current = DeepgramSTTSettings(
|
||||||
model="nova-3-general",
|
model="nova-3-general",
|
||||||
language="en",
|
language="en",
|
||||||
encoding="linear16",
|
|
||||||
channels=1,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
raw = {"model": "nova-2"}
|
raw = {"model": "nova-2"}
|
||||||
@@ -474,16 +467,13 @@ class TestDeepgramSageMakerSTTSettings:
|
|||||||
store = DeepgramSageMakerSTTSettings(
|
store = DeepgramSageMakerSTTSettings(
|
||||||
model="nova-3",
|
model="nova-3",
|
||||||
language="en",
|
language="en",
|
||||||
encoding="linear16",
|
|
||||||
channels=1,
|
|
||||||
punctuate=True,
|
|
||||||
)
|
)
|
||||||
delta = DeepgramSageMakerSTTSettings(punctuate=False)
|
delta = DeepgramSageMakerSTTSettings(model="nova-2")
|
||||||
changed = store.apply_update(delta)
|
changed = store.apply_update(delta)
|
||||||
|
|
||||||
assert store.punctuate is False
|
assert store.model == "nova-2"
|
||||||
assert store.encoding == "linear16"
|
assert store.language == "en"
|
||||||
assert "punctuate" in changed
|
assert "model" in changed
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
@@ -499,17 +489,17 @@ class TestDeepgramSTTSettingsExtraSync:
|
|||||||
return DeepgramSTTService(api_key="test-key", sample_rate=16000, **kwargs)
|
return DeepgramSTTService(api_key="test-key", sample_rate=16000, **kwargs)
|
||||||
|
|
||||||
def test_extra_synced_to_declared_field_at_init(self):
|
def test_extra_synced_to_declared_field_at_init(self):
|
||||||
"""If LiveOptions has unknown params in _extra, they can be synced if they match fields."""
|
"""LiveOptions params that match declared fields are synced at init."""
|
||||||
from pipecat.services.deepgram.stt import LiveOptions
|
from pipecat.services.deepgram.stt import LiveOptions
|
||||||
|
|
||||||
# Use **kwargs to pass undeclared params
|
live_options = LiveOptions(numerals=True)
|
||||||
live_options = LiveOptions(numerals=True) # 'numerals' goes into _extra
|
|
||||||
|
|
||||||
svc = self._make_service(live_options=live_options)
|
svc = self._make_service(live_options=live_options)
|
||||||
|
|
||||||
# 'numerals' doesn't match a declared DeepgramSTTSettings field,
|
# 'numerals' is a declared DeepgramSTTSettings field,
|
||||||
# so it should stay in extra
|
# so it should be promoted from extra to the declared field
|
||||||
assert svc._settings.extra["numerals"] is True
|
assert svc._settings.numerals is True
|
||||||
|
assert "numerals" not in svc._settings.extra
|
||||||
|
|
||||||
def test_declared_field_from_live_options(self):
|
def test_declared_field_from_live_options(self):
|
||||||
"""LiveOptions fields that match DeepgramSTTSettings fields are applied."""
|
"""LiveOptions fields that match DeepgramSTTSettings fields are applied."""
|
||||||
@@ -532,7 +522,7 @@ class TestDeepgramSTTSettingsExtraSync:
|
|||||||
raw_dict = {
|
raw_dict = {
|
||||||
"diarize": True, # matches declared field
|
"diarize": True, # matches declared field
|
||||||
"punctuate": False, # matches declared field
|
"punctuate": False, # matches declared field
|
||||||
"numerals": True, # doesn't match - stays in extra
|
"custom_param": "value", # doesn't match - stays in extra
|
||||||
}
|
}
|
||||||
|
|
||||||
delta = DeepgramSTTSettings.from_mapping(raw_dict)
|
delta = DeepgramSTTSettings.from_mapping(raw_dict)
|
||||||
@@ -541,7 +531,7 @@ class TestDeepgramSTTSettingsExtraSync:
|
|||||||
assert delta.diarize is True
|
assert delta.diarize is True
|
||||||
assert delta.punctuate is False
|
assert delta.punctuate is False
|
||||||
# Unknown stays in extra
|
# Unknown stays in extra
|
||||||
assert delta.extra["numerals"] is True
|
assert delta.extra["custom_param"] == "value"
|
||||||
|
|
||||||
# Now simulate syncing (though from_mapping already routes correctly)
|
# Now simulate syncing (though from_mapping already routes correctly)
|
||||||
delta._sync_extra_to_fields()
|
delta._sync_extra_to_fields()
|
||||||
@@ -549,7 +539,7 @@ class TestDeepgramSTTSettingsExtraSync:
|
|||||||
# Still the same - from_mapping already put them in the right place
|
# Still the same - from_mapping already put them in the right place
|
||||||
assert delta.diarize is True
|
assert delta.diarize is True
|
||||||
assert delta.punctuate is False
|
assert delta.punctuate is False
|
||||||
assert delta.extra["numerals"] is True
|
assert delta.extra["custom_param"] == "value"
|
||||||
|
|
||||||
def test_sync_promotes_extra_to_field_when_not_given(self):
|
def test_sync_promotes_extra_to_field_when_not_given(self):
|
||||||
"""_sync_extra_to_fields promotes extra dict entries to declared fields."""
|
"""_sync_extra_to_fields promotes extra dict entries to declared fields."""
|
||||||
@@ -611,16 +601,17 @@ class TestDeepgramSTTSettingsExtraSync:
|
|||||||
"""Unknown params (not matching fields) stay in extra and get forwarded."""
|
"""Unknown params (not matching fields) stay in extra and get forwarded."""
|
||||||
from pipecat.services.deepgram.stt import LiveOptions
|
from pipecat.services.deepgram.stt import LiveOptions
|
||||||
|
|
||||||
# numerals isn't a declared field in DeepgramSTTSettings
|
# 'numerals' is now a declared field; 'custom_param' is not
|
||||||
live_options = LiveOptions(numerals=True, custom_param="test")
|
live_options = LiveOptions(numerals=True, custom_param="test")
|
||||||
|
|
||||||
svc = self._make_service(live_options=live_options)
|
svc = self._make_service(live_options=live_options)
|
||||||
|
|
||||||
# Should be in extra
|
# 'numerals' is a declared field, so it should be promoted
|
||||||
assert svc._settings.extra["numerals"] is True
|
assert svc._settings.numerals is True
|
||||||
|
# 'custom_param' is unknown, so it stays in extra
|
||||||
assert svc._settings.extra["custom_param"] == "test"
|
assert svc._settings.extra["custom_param"] == "test"
|
||||||
|
|
||||||
# And forwarded to kwargs
|
# Both forwarded to kwargs
|
||||||
kwargs = svc._build_connect_kwargs()
|
kwargs = svc._build_connect_kwargs()
|
||||||
assert kwargs["numerals"] == "true"
|
assert kwargs["numerals"] == "true"
|
||||||
assert kwargs["custom_param"] == "test"
|
assert kwargs["custom_param"] == "test"
|
||||||
|
|||||||
Reference in New Issue
Block a user