Update InputParams to languages: support str or List of Languages
This commit is contained in:
@@ -45,7 +45,7 @@ async def main():
|
|||||||
)
|
)
|
||||||
|
|
||||||
stt = GoogleSTTService(
|
stt = GoogleSTTService(
|
||||||
params=GoogleSTTService.InputParams(language=Language.EN_US),
|
params=GoogleSTTService.InputParams(languages=Language.EN_US),
|
||||||
)
|
)
|
||||||
|
|
||||||
tts = GoogleTTSService(
|
tts = GoogleTTSService(
|
||||||
|
|||||||
@@ -14,11 +14,11 @@ import os
|
|||||||
os.environ["GRPC_ENABLE_FORK_SUPPORT"] = "false"
|
os.environ["GRPC_ENABLE_FORK_SUPPORT"] = "false"
|
||||||
|
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, AsyncGenerator, Dict, List, Literal, Optional
|
from typing import Any, AsyncGenerator, Dict, List, Literal, Optional, Union
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field, field_validator
|
||||||
|
|
||||||
from pipecat.frames.frames import (
|
from pipecat.frames.frames import (
|
||||||
AudioRawFrame,
|
AudioRawFrame,
|
||||||
@@ -1425,7 +1425,7 @@ class GoogleSTTService(STTService):
|
|||||||
"""Configuration parameters for Google Speech-to-Text.
|
"""Configuration parameters for Google Speech-to-Text.
|
||||||
|
|
||||||
Attributes:
|
Attributes:
|
||||||
language: Recognition language (defaults to US English).
|
languages: Single language or list of recognition languages. First language is primary.
|
||||||
model: Speech recognition model to use.
|
model: Speech recognition model to use.
|
||||||
use_separate_recognition_per_channel: Process each audio channel separately.
|
use_separate_recognition_per_channel: Process each audio channel separately.
|
||||||
enable_automatic_punctuation: Add punctuation to transcripts.
|
enable_automatic_punctuation: Add punctuation to transcripts.
|
||||||
@@ -1438,7 +1438,7 @@ class GoogleSTTService(STTService):
|
|||||||
enable_voice_activity_events: Detect voice activity in audio.
|
enable_voice_activity_events: Detect voice activity in audio.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
language: Optional[Language] = Language.EN_US
|
languages: Union[Language, List[Language]] = Field(default_factory=lambda: [Language.EN_US])
|
||||||
model: Optional[str] = "latest_long"
|
model: Optional[str] = "latest_long"
|
||||||
use_separate_recognition_per_channel: Optional[bool] = False
|
use_separate_recognition_per_channel: Optional[bool] = False
|
||||||
enable_automatic_punctuation: Optional[bool] = True
|
enable_automatic_punctuation: Optional[bool] = True
|
||||||
@@ -1450,6 +1450,19 @@ class GoogleSTTService(STTService):
|
|||||||
enable_interim_results: Optional[bool] = True
|
enable_interim_results: Optional[bool] = True
|
||||||
enable_voice_activity_events: Optional[bool] = False
|
enable_voice_activity_events: Optional[bool] = False
|
||||||
|
|
||||||
|
@field_validator("languages", mode="before")
|
||||||
|
@classmethod
|
||||||
|
def validate_languages(cls, v) -> List[Language]:
|
||||||
|
if isinstance(v, Language):
|
||||||
|
return [v]
|
||||||
|
return v
|
||||||
|
|
||||||
|
@property
|
||||||
|
def language_list(self) -> List[Language]:
|
||||||
|
"""Get languages as a guaranteed list."""
|
||||||
|
assert isinstance(self.languages, list)
|
||||||
|
return self.languages
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
@@ -1506,9 +1519,9 @@ class GoogleSTTService(STTService):
|
|||||||
self._client = speech_v2.SpeechAsyncClient(credentials=creds, client_options=client_options)
|
self._client = speech_v2.SpeechAsyncClient(credentials=creds, client_options=client_options)
|
||||||
|
|
||||||
self._settings = {
|
self._settings = {
|
||||||
"language_code": self.language_to_service_language(params.language)
|
"language_codes": [
|
||||||
if params.language
|
self.language_to_service_language(lang) for lang in params.language_list
|
||||||
else "en-US",
|
],
|
||||||
"model": params.model,
|
"model": params.model,
|
||||||
"use_separate_recognition_per_channel": params.use_separate_recognition_per_channel,
|
"use_separate_recognition_per_channel": params.use_separate_recognition_per_channel,
|
||||||
"enable_automatic_punctuation": params.enable_automatic_punctuation,
|
"enable_automatic_punctuation": params.enable_automatic_punctuation,
|
||||||
@@ -1521,22 +1534,30 @@ class GoogleSTTService(STTService):
|
|||||||
"enable_voice_activity_events": params.enable_voice_activity_events,
|
"enable_voice_activity_events": params.enable_voice_activity_events,
|
||||||
}
|
}
|
||||||
|
|
||||||
def language_to_service_language(self, language: Language) -> Optional[str]:
|
def language_to_service_language(self, language: Language | List[Language]) -> str | List[str]:
|
||||||
"""Convert Language enum to Google STT language code.
|
"""Convert Language enum(s) to Google STT language code(s).
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
language: Language enum value.
|
language: Single Language enum or list of Language enums.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
str: Google STT language code.
|
str | List[str]: Google STT language code(s).
|
||||||
"""
|
"""
|
||||||
return language_to_google_stt_language(language)
|
if isinstance(language, list):
|
||||||
|
return [language_to_google_stt_language(lang) or "en-US" for lang in language]
|
||||||
|
return language_to_google_stt_language(language) or "en-US"
|
||||||
|
|
||||||
async def set_language(self, language: Language):
|
async def set_languages(self, languages: List[Language]):
|
||||||
"""Update the service's recognition language."""
|
"""Update the service's recognition languages.
|
||||||
logger.info(f"Switching STT language to: [{language}]")
|
|
||||||
self._settings["language_code"] = self.language_to_service_language(language)
|
Args:
|
||||||
# Recreate stream with new language
|
languages: List of languages for recognition. First language is primary.
|
||||||
|
"""
|
||||||
|
logger.info(f"Switching STT languages to: {languages}")
|
||||||
|
self._settings["language_codes"] = [
|
||||||
|
self.language_to_service_language(lang) for lang in languages
|
||||||
|
]
|
||||||
|
# Recreate stream with new languages
|
||||||
if self._streaming_task:
|
if self._streaming_task:
|
||||||
await self._disconnect()
|
await self._disconnect()
|
||||||
await self._connect()
|
await self._connect()
|
||||||
@@ -1565,7 +1586,7 @@ class GoogleSTTService(STTService):
|
|||||||
async def update_options(
|
async def update_options(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
language: Optional[Language] = None,
|
languages: Optional[List[Language]] = None,
|
||||||
model: Optional[str] = None,
|
model: Optional[str] = None,
|
||||||
enable_automatic_punctuation: Optional[bool] = None,
|
enable_automatic_punctuation: Optional[bool] = None,
|
||||||
enable_spoken_punctuation: Optional[bool] = None,
|
enable_spoken_punctuation: Optional[bool] = None,
|
||||||
@@ -1580,7 +1601,7 @@ class GoogleSTTService(STTService):
|
|||||||
"""Update service options dynamically.
|
"""Update service options dynamically.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
language: New recognition language.
|
languages: New list of recongition languages.
|
||||||
model: New recognition model.
|
model: New recognition model.
|
||||||
enable_automatic_punctuation: Enable/disable automatic punctuation.
|
enable_automatic_punctuation: Enable/disable automatic punctuation.
|
||||||
enable_spoken_punctuation: Enable/disable spoken punctuation.
|
enable_spoken_punctuation: Enable/disable spoken punctuation.
|
||||||
@@ -1599,9 +1620,11 @@ class GoogleSTTService(STTService):
|
|||||||
needs_reconnect = False
|
needs_reconnect = False
|
||||||
|
|
||||||
# Update settings with new values
|
# Update settings with new values
|
||||||
if language is not None:
|
if languages is not None:
|
||||||
logger.debug(f"Updating language to: {language}")
|
logger.debug(f"Updating language to: {languages}")
|
||||||
self._settings["language_code"] = self.language_to_service_language(language)
|
self._settings["language_codes"] = [
|
||||||
|
self.language_to_service_language(lang) for lang in languages
|
||||||
|
]
|
||||||
needs_reconnect = True
|
needs_reconnect = True
|
||||||
|
|
||||||
if model is not None:
|
if model is not None:
|
||||||
@@ -1672,7 +1695,7 @@ class GoogleSTTService(STTService):
|
|||||||
sample_rate_hertz=self.sample_rate,
|
sample_rate_hertz=self.sample_rate,
|
||||||
audio_channel_count=1,
|
audio_channel_count=1,
|
||||||
),
|
),
|
||||||
language_codes=[self._settings["language_code"]],
|
language_codes=self._settings["language_codes"],
|
||||||
model=self._settings["model"],
|
model=self._settings["model"],
|
||||||
features=cloud_speech.RecognitionFeatures(
|
features=cloud_speech.RecognitionFeatures(
|
||||||
enable_automatic_punctuation=self._settings["enable_automatic_punctuation"],
|
enable_automatic_punctuation=self._settings["enable_automatic_punctuation"],
|
||||||
@@ -1775,16 +1798,17 @@ class GoogleSTTService(STTService):
|
|||||||
if not transcript:
|
if not transcript:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
|
# Use the primary language (first in the list)
|
||||||
|
primary_language = self._settings["language_codes"][0]
|
||||||
|
|
||||||
if result.is_final:
|
if result.is_final:
|
||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
TranscriptionFrame(
|
TranscriptionFrame(transcript, "", time_now_iso8601(), primary_language)
|
||||||
transcript, "", time_now_iso8601(), self._settings["language_code"]
|
|
||||||
)
|
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
await self.push_frame(
|
await self.push_frame(
|
||||||
InterimTranscriptionFrame(
|
InterimTranscriptionFrame(
|
||||||
transcript, "", time_now_iso8601(), self._settings["language_code"]
|
transcript, "", time_now_iso8601(), primary_language
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user