Add delay_in_frames and language support

This commit is contained in:
Mark Backman
2026-01-29 10:59:04 -05:00
parent 6ab12626d6
commit 31c7fbc5ba
4 changed files with 92 additions and 7 deletions

View File

@@ -1 +1,3 @@
- `GradiumSTTService` now flushes pending transcriptions when VAD detects the user stopped speaking, improving response latency. - Updates to `GradiumSTTService`:
- Now flushes pending transcriptions when VAD detects the user stopped speaking, improving response latency.
- `GradiumSTTService` now supports `InputParams` for configuring `language` and `delay_in_frames` settings.

View File

@@ -26,6 +26,7 @@ from pipecat.runner.utils import create_transport
from pipecat.services.gradium.stt import GradiumSTTService from pipecat.services.gradium.stt import GradiumSTTService
from pipecat.services.gradium.tts import GradiumTTSService from pipecat.services.gradium.tts import GradiumTTSService
from pipecat.services.openai.llm import OpenAILLMService from pipecat.services.openai.llm import OpenAILLMService
from pipecat.transcriptions.language import Language
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
@@ -62,6 +63,9 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
stt = GradiumSTTService( stt = GradiumSTTService(
api_key=os.getenv("GRADIUM_API_KEY"), api_key=os.getenv("GRADIUM_API_KEY"),
api_endpoint_base_url="wss://us.api.gradium.ai/api/speech/asr", api_endpoint_base_url="wss://us.api.gradium.ai/api/speech/asr",
params=GradiumSTTService.InputParams(
language=Language.EN,
),
) )
tts = GradiumTTSService( tts = GradiumTTSService(

View File

@@ -17,6 +17,7 @@ from pipecat.processors.frame_processor import FrameDirection, FrameProcessor
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.gradium.stt import GradiumSTTService from pipecat.services.gradium.stt import GradiumSTTService
from pipecat.transcriptions.language import Language
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
@@ -51,6 +52,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
stt = GradiumSTTService( stt = GradiumSTTService(
api_key=os.getenv("GRADIUM_API_KEY"), api_key=os.getenv("GRADIUM_API_KEY"),
api_endpoint_base_url="wss://us.api.gradium.ai/api/speech/asr", api_endpoint_base_url="wss://us.api.gradium.ai/api/speech/asr",
params=GradiumSTTService.InputParams(language=Language.EN, delay_in_frames=8),
) )
tl = TranscriptionLogger() tl = TranscriptionLogger()

View File

@@ -12,9 +12,10 @@ WebSocket API for streaming audio transcription.
import base64 import base64
import json import json
from typing import AsyncGenerator from typing import AsyncGenerator, Optional
from loguru import logger from loguru import logger
from pydantic import BaseModel
from pipecat.frames.frames import ( from pipecat.frames.frames import (
CancelFrame, CancelFrame,
@@ -27,7 +28,7 @@ from pipecat.frames.frames import (
) )
from pipecat.processors.frame_processor import FrameDirection from pipecat.processors.frame_processor import FrameDirection
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, resolve_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
@@ -42,6 +43,26 @@ except ModuleNotFoundError as e:
SAMPLE_RATE = 24000 SAMPLE_RATE = 24000
def language_to_gradium_language(language: Language) -> Optional[str]:
"""Convert a Language enum to Gradium's language code format.
Args:
language: The Language enum value to convert.
Returns:
The Gradium language code string or None if not supported.
"""
LANGUAGE_MAP = {
Language.DE: "de",
Language.EN: "en",
Language.ES: "es",
Language.FR: "fr",
Language.PT: "pt",
}
return resolve_language(language, LANGUAGE_MAP, use_base_code=True)
class GradiumSTTService(WebsocketSTTService): class GradiumSTTService(WebsocketSTTService):
"""Gradium real-time speech-to-text service. """Gradium real-time speech-to-text service.
@@ -50,12 +71,29 @@ class GradiumSTTService(WebsocketSTTService):
for audio processing and connection management. for audio processing and connection management.
""" """
class InputParams(BaseModel):
"""Configuration parameters for Gradium STT API.
Parameters:
language: Expected language of the audio (e.g., "en", "es", "fr").
This helps ground the model to a specific language and improve
transcription quality.
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.
Default is 10 (800ms). Lower values like 7-8 give faster response.
"""
language: Optional[Language] = None
delay_in_frames: Optional[int] = None
def __init__( def __init__(
self, self,
*, *,
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",
json_config: str | None = None, params: Optional[InputParams] = None,
json_config: Optional[str] = None,
**kwargs, **kwargs,
): ):
"""Initialize the Gradium STT service. """Initialize the Gradium STT service.
@@ -63,14 +101,29 @@ 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.
params: Configuration parameters for language and delay settings.
json_config: Optional JSON configuration string for additional model settings. json_config: Optional JSON configuration string for additional model settings.
.. deprecated:: 0.0.101
Use `params` instead for type-safe configuration.
**kwargs: Additional arguments passed to parent STTService class. **kwargs: Additional arguments passed to parent STTService class.
""" """
super().__init__(sample_rate=SAMPLE_RATE, **kwargs) super().__init__(sample_rate=SAMPLE_RATE, **kwargs)
if json_config is not None:
import warnings
warnings.warn(
"Parameter 'json_config' is deprecated and will be removed in a future version, use 'params' instead.",
DeprecationWarning,
stacklevel=2,
)
self._api_key = api_key self._api_key = api_key
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._params = params or GradiumSTTService.InputParams()
self._json_config = json_config self._json_config = json_config
self._receive_task = None self._receive_task = None
@@ -92,6 +145,17 @@ class GradiumSTTService(WebsocketSTTService):
""" """
return True return True
async def set_language(self, language: Language):
"""Set the recognition language and reconnect.
Args:
language: The language to use for speech recognition.
"""
logger.info(f"Switching STT language to: [{language}]")
self._params.language = language
await self._disconnect()
await self._connect()
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
"""Start the speech-to-text service. """Start the speech-to-text service.
@@ -226,8 +290,18 @@ class GradiumSTTService(WebsocketSTTService):
"type": "setup", "type": "setup",
"input_format": "pcm", "input_format": "pcm",
} }
if self._json_config is not None: # Build json_config: start with deprecated json_config, then override with params
setup_msg["json_config"] = self._json_config json_config = {}
if self._json_config:
json_config = json.loads(self._json_config)
if self._params.language:
gradium_language = language_to_gradium_language(self._params.language)
if gradium_language:
json_config["language"] = gradium_language
if self._params.delay_in_frames:
json_config["delay_in_frames"] = self._params.delay_in_frames
if json_config:
setup_msg["json_config"] = json_config
await self._websocket.send(json.dumps(setup_msg)) await self._websocket.send(json.dumps(setup_msg))
ready_msg = await self._websocket.recv() ready_msg = await self._websocket.recv()
ready_msg = json.loads(ready_msg) ready_msg = json.loads(ready_msg)
@@ -239,7 +313,10 @@ class GradiumSTTService(WebsocketSTTService):
# Store delay_in_frames and frame_size for silence flushing # Store delay_in_frames and frame_size for silence flushing
self._delay_in_frames = ready_msg.get("delay_in_frames", 0) self._delay_in_frames = ready_msg.get("delay_in_frames", 0)
self._frame_size = ready_msg.get("frame_size", 1920) self._frame_size = ready_msg.get("frame_size", 1920)
logger.debug(f"Connected to Gradium STT") logger.debug(
f"Connected to Gradium STT (delay_in_frames={self._delay_in_frames}, "
f"frame_size={self._frame_size})"
)
except Exception as e: except Exception as e:
await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e) await self.push_error(error_msg=f"Unknown error occurred: {e}", exception=e)