Update GladiaSTTService to use the Gladia V2 API

This commit is contained in:
Mark Backman
2024-10-20 10:32:53 -04:00
parent b6b1ef0a40
commit 46927805bc
3 changed files with 108 additions and 121 deletions

View File

@@ -15,6 +15,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
- Added a foundational example for Gladia transcription: - Added a foundational example for Gladia transcription:
`13c-gladia-transcription.py` `13c-gladia-transcription.py`
### Changed
- Updated `GladiaSTTService` to use the V2 API.
### Fixed ### Fixed
- Fixed `enable_usage_metrics` to control LLM/TTS usage metrics separately - Fixed `enable_usage_metrics` to control LLM/TTS usage metrics separately

View File

@@ -5,12 +5,16 @@
# #
import asyncio import asyncio
import aiohttp
import os import os
import sys import sys
import aiohttp
from dotenv import load_dotenv
from loguru import logger
from runner import configure
from pipecat.audio.vad.silero import SileroVADAnalyzer from pipecat.audio.vad.silero import SileroVADAnalyzer
from pipecat.frames.frames import LLMMessagesFrame from pipecat.frames.frames import EndFrame, LLMMessagesFrame
from pipecat.pipeline.pipeline import Pipeline from pipecat.pipeline.pipeline import Pipeline
from pipecat.pipeline.runner import PipelineRunner from pipecat.pipeline.runner import PipelineRunner
from pipecat.pipeline.task import PipelineParams, PipelineTask from pipecat.pipeline.task import PipelineParams, PipelineTask
@@ -20,12 +24,6 @@ from pipecat.services.gladia import GladiaSTTService
from pipecat.services.openai import OpenAILLMService from pipecat.services.openai import OpenAILLMService
from pipecat.transports.services.daily import DailyParams, DailyTransport from pipecat.transports.services.daily import DailyParams, DailyTransport
from runner import configure
from loguru import logger
from dotenv import load_dotenv
load_dotenv(override=True) load_dotenv(override=True)
logger.remove(0) logger.remove(0)
@@ -90,6 +88,11 @@ async def main():
messages.append({"role": "system", "content": "Please introduce yourself to the user."}) messages.append({"role": "system", "content": "Please introduce yourself to the user."})
await task.queue_frames([LLMMessagesFrame(messages)]) await task.queue_frames([LLMMessagesFrame(messages)])
# Register an event handler to exit the application when the user leaves.
@transport.event_handler("on_participant_left")
async def on_participant_left(transport, participant, reason):
await task.queue_frame(EndFrame())
runner = PipelineRunner() runner = PipelineRunner()
await runner.run(task) await runner.run(task)

View File

@@ -8,6 +8,7 @@ import base64
import json import json
from typing import AsyncGenerator, Optional from typing import AsyncGenerator, Optional
import aiohttp
from loguru import logger from loguru import logger
from pydantic.main import BaseModel from pydantic.main import BaseModel
@@ -23,7 +24,6 @@ from pipecat.services.ai_services 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
# See .env.example for Gladia configuration needed
try: try:
import websockets import websockets
except ModuleNotFoundError as e: except ModuleNotFoundError as e:
@@ -38,15 +38,16 @@ class GladiaSTTService(STTService):
class InputParams(BaseModel): class InputParams(BaseModel):
sample_rate: Optional[int] = 16000 sample_rate: Optional[int] = 16000
language: Optional[Language] = Language.EN language: Optional[Language] = Language.EN
transcription_hint: Optional[str] = None endpointing: Optional[float] = 0.2
endpointing: Optional[int] = 200 maximum_duration_without_endpointing: Optional[int] = 10
prosody: Optional[bool] = None audio_enhancer: Optional[bool] = None
words_accurate_timestamps: Optional[bool] = None
def __init__( def __init__(
self, self,
*, *,
api_key: str, api_key: str,
url: str = "wss://api.gladia.io/audio/text/audio-transcription", url: str = "https://api.gladia.io/v2/live",
confidence: float = 0.5, confidence: float = 0.5,
params: InputParams = InputParams(), params: InputParams = InputParams(),
**kwargs, **kwargs,
@@ -56,101 +57,82 @@ class GladiaSTTService(STTService):
self._api_key = api_key self._api_key = api_key
self._url = url self._url = url
self._settings = { self._settings = {
"encoding": "wav/pcm",
"bit_depth": 16,
"sample_rate": params.sample_rate, "sample_rate": params.sample_rate,
"language": self.language_to_service_language(params.language) "channels": 1,
if params.language "language_config": {
else Language.EN, "languages": [self.language_to_service_language(params.language)]
"transcription_hint": params.transcription_hint, if params.language
else [],
"code_switching": False,
},
"endpointing": params.endpointing, "endpointing": params.endpointing,
"prosody": params.prosody, "maximum_duration_without_endpointing": params.maximum_duration_without_endpointing,
"pre_processing": {
"audio_enhancer": params.audio_enhancer,
},
"realtime_processing": {
"words_accurate_timestamps": params.words_accurate_timestamps,
},
} }
self._confidence = confidence self._confidence = confidence
def language_to_service_language(self, language: Language) -> str | None: def language_to_service_language(self, language: Language) -> str | None:
match language: language_map = {
case Language.BG: Language.BG: "bg",
return "bulgarian" Language.CA: "ca",
case Language.CA: Language.ZH: "zh",
return "catalan" Language.CS: "cs",
case Language.ZH: Language.DA: "da",
return "chinese" Language.NL: "nl",
case Language.CS: Language.EN: "en",
return "czech" Language.EN_US: "en",
case Language.DA: Language.EN_AU: "en",
return "danish" Language.EN_GB: "en",
case Language.NL: Language.EN_NZ: "en",
return "dutch" Language.EN_IN: "en",
case ( Language.ET: "et",
Language.EN Language.FI: "fi",
| Language.EN_US Language.FR: "fr",
| Language.EN_AU Language.FR_CA: "fr",
| Language.EN_GB Language.DE: "de",
| Language.EN_NZ Language.DE_CH: "de",
| Language.EN_IN Language.EL: "el",
): Language.HI: "hi",
return "english" Language.HU: "hu",
case Language.ET: Language.ID: "id",
return "estonian" Language.IT: "it",
case Language.FI: Language.JA: "ja",
return "finnish" Language.KO: "ko",
case Language.FR | Language.FR_CA: Language.LV: "lv",
return "french" Language.LT: "lt",
case Language.DE | Language.DE_CH: Language.MS: "ms",
return "german" Language.NO: "no",
case Language.EL: Language.PL: "pl",
return "greek" Language.PT: "pt",
case Language.HI: Language.PT_BR: "pt",
return "hindi" Language.RO: "ro",
case Language.HU: Language.RU: "ru",
return "hungarian" Language.SK: "sk",
case Language.ID: Language.ES: "es",
return "indonesian" Language.SV: "sv",
case Language.IT: Language.TH: "th",
return "italian" Language.TR: "tr",
case Language.JA: Language.UK: "uk",
return "japanese" Language.VI: "vi",
case Language.KO: }
return "korean" return language_map.get(language)
case Language.LV:
return "latvian"
case Language.LT:
return "lithuanian"
case Language.MS:
return "malay"
case Language.NO:
return "norwegian"
case Language.PL:
return "polish"
case Language.PT | Language.PT_BR:
return "portuguese"
case Language.RO:
return "romanian"
case Language.RU:
return "russian"
case Language.SK:
return "slovak"
case Language.ES:
return "spanish"
case Language.SV:
return "slovenian"
case Language.TH:
return "thai"
case Language.TR:
return "turkish"
case Language.UK:
return "ukrainian"
case Language.VI:
return "vietnamese"
return None
async def start(self, frame: StartFrame): async def start(self, frame: StartFrame):
await super().start(frame) await super().start(frame)
self._websocket = await websockets.connect(self._url) response = await self._setup_gladia()
self._websocket = await websockets.connect(response["url"])
self._receive_task = self.get_event_loop().create_task(self._receive_task_handler()) self._receive_task = self.get_event_loop().create_task(self._receive_task_handler())
await self._setup_gladia()
async def stop(self, frame: EndFrame): async def stop(self, frame: EndFrame):
await super().stop(frame) await super().stop(frame)
await self._send_stop_recording()
await self._websocket.close() await self._websocket.close()
async def cancel(self, frame: CancelFrame): async def cancel(self, frame: CancelFrame):
@@ -164,39 +146,37 @@ class GladiaSTTService(STTService):
yield None yield None
async def _setup_gladia(self): async def _setup_gladia(self):
configuration = { async with aiohttp.ClientSession() as session:
"x_gladia_key": self._api_key, async with session.post(
"encoding": "WAV/PCM", self._url,
"model_type": "fast", headers={"X-Gladia-Key": self._api_key, "Content-Type": "application/json"},
"language_behaviour": "manual", json=self._settings,
"sample_rate": self._settings["sample_rate"], ) as response:
"language": self._settings["language"], if response.ok:
"transcription_hint": self._settings["transcription_hint"], return await response.json()
"endpointing": self._settings["endpointing"], else:
"prosody": self._settings["prosody"], logger.error(
} f"Gladia error: {response.status}: {response.text or response.reason}"
)
await self._websocket.send(json.dumps(configuration)) raise Exception(f"Failed to initialize Gladia session: {response.status}")
async def _send_audio(self, audio: bytes): async def _send_audio(self, audio: bytes):
message = {"frames": base64.b64encode(audio).decode("utf-8")} data = base64.b64encode(audio).decode("utf-8")
message = {"type": "audio_chunk", "data": {"chunk": data}}
await self._websocket.send(json.dumps(message)) await self._websocket.send(json.dumps(message))
async def _send_stop_recording(self):
await self._websocket.send(json.dumps({"type": "stop_recording"}))
async def _receive_task_handler(self): async def _receive_task_handler(self):
async for message in self._websocket: async for message in self._websocket:
utterance = json.loads(message) content = json.loads(message)
if not utterance: if content["type"] == "transcript":
continue utterance = content["data"]["utterance"]
confidence = utterance.get("confidence", 0)
if "error" in utterance: transcript = utterance["text"]
message = utterance["message"]
logger.error(f"Gladia error: {message}")
elif "confidence" in utterance:
type = utterance["type"]
confidence = utterance["confidence"]
transcript = utterance["transcription"]
if confidence >= self._confidence: if confidence >= self._confidence:
if type == "final": if content["data"]["is_final"]:
await self.push_frame( await self.push_frame(
TranscriptionFrame(transcript, "", time_now_iso8601()) TranscriptionFrame(transcript, "", time_now_iso8601())
) )