Merge pull request #1190 from pipecat-ai/mb/languages-hosted-whisper
Add language support to OpenAI and Groq hosted Whisper
This commit is contained in:
@@ -10,6 +10,7 @@ from loguru import logger
|
|||||||
|
|
||||||
from pipecat.frames.frames import ErrorFrame, Frame, TranscriptionFrame
|
from pipecat.frames.frames import ErrorFrame, Frame, TranscriptionFrame
|
||||||
from pipecat.services.ai_services import SegmentedSTTService
|
from pipecat.services.ai_services import SegmentedSTTService
|
||||||
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -23,6 +24,82 @@ except ModuleNotFoundError as e:
|
|||||||
raise Exception(f"Missing module: {e}")
|
raise Exception(f"Missing module: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
def language_to_whisper_language(language: Language) -> Optional[str]:
|
||||||
|
"""Language support for Whisper API.
|
||||||
|
|
||||||
|
Docs: https://platform.openai.com/docs/guides/speech-to-text#supported-languages
|
||||||
|
"""
|
||||||
|
BASE_LANGUAGES = {
|
||||||
|
Language.AF: "af",
|
||||||
|
Language.AR: "ar",
|
||||||
|
Language.HY: "hy",
|
||||||
|
Language.AZ: "az",
|
||||||
|
Language.BE: "be",
|
||||||
|
Language.BS: "bs",
|
||||||
|
Language.BG: "bg",
|
||||||
|
Language.CA: "ca",
|
||||||
|
Language.ZH: "zh",
|
||||||
|
Language.HR: "hr",
|
||||||
|
Language.CS: "cs",
|
||||||
|
Language.DA: "da",
|
||||||
|
Language.NL: "nl",
|
||||||
|
Language.EN: "en",
|
||||||
|
Language.ET: "et",
|
||||||
|
Language.FI: "fi",
|
||||||
|
Language.FR: "fr",
|
||||||
|
Language.GL: "gl",
|
||||||
|
Language.DE: "de",
|
||||||
|
Language.EL: "el",
|
||||||
|
Language.HE: "he",
|
||||||
|
Language.HI: "hi",
|
||||||
|
Language.HU: "hu",
|
||||||
|
Language.IS: "is",
|
||||||
|
Language.ID: "id",
|
||||||
|
Language.IT: "it",
|
||||||
|
Language.JA: "ja",
|
||||||
|
Language.KN: "kn",
|
||||||
|
Language.KK: "kk",
|
||||||
|
Language.KO: "ko",
|
||||||
|
Language.LV: "lv",
|
||||||
|
Language.LT: "lt",
|
||||||
|
Language.MK: "mk",
|
||||||
|
Language.MS: "ms",
|
||||||
|
Language.MR: "mr",
|
||||||
|
Language.MI: "mi",
|
||||||
|
Language.NE: "ne",
|
||||||
|
Language.NO: "no",
|
||||||
|
Language.FA: "fa",
|
||||||
|
Language.PL: "pl",
|
||||||
|
Language.PT: "pt",
|
||||||
|
Language.RO: "ro",
|
||||||
|
Language.RU: "ru",
|
||||||
|
Language.SR: "sr",
|
||||||
|
Language.SK: "sk",
|
||||||
|
Language.SL: "sl",
|
||||||
|
Language.ES: "es",
|
||||||
|
Language.SW: "sw",
|
||||||
|
Language.SV: "sv",
|
||||||
|
Language.TL: "tl",
|
||||||
|
Language.TA: "ta",
|
||||||
|
Language.TH: "th",
|
||||||
|
Language.TR: "tr",
|
||||||
|
Language.UK: "uk",
|
||||||
|
Language.UR: "ur",
|
||||||
|
Language.VI: "vi",
|
||||||
|
Language.CY: "cy",
|
||||||
|
}
|
||||||
|
|
||||||
|
result = BASE_LANGUAGES.get(language)
|
||||||
|
|
||||||
|
# If not found in base languages, try to find the base language from a variant
|
||||||
|
if not result:
|
||||||
|
lang_str = str(language.value)
|
||||||
|
base_code = lang_str.split("-")[0].lower()
|
||||||
|
result = base_code if base_code in BASE_LANGUAGES.values() else None
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
class BaseWhisperSTTService(SegmentedSTTService):
|
class BaseWhisperSTTService(SegmentedSTTService):
|
||||||
"""Base class for Whisper-based speech-to-text services.
|
"""Base class for Whisper-based speech-to-text services.
|
||||||
|
|
||||||
@@ -33,6 +110,7 @@ class BaseWhisperSTTService(SegmentedSTTService):
|
|||||||
model: Name of the Whisper model to use.
|
model: Name of the Whisper model to use.
|
||||||
api_key: Service API key. Defaults to None.
|
api_key: Service API key. Defaults to None.
|
||||||
base_url: Service API base URL. Defaults to None.
|
base_url: Service API base URL. Defaults to None.
|
||||||
|
language: Language of the audio input. Defaults to English.
|
||||||
**kwargs: Additional arguments passed to SegmentedSTTService.
|
**kwargs: Additional arguments passed to SegmentedSTTService.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -42,11 +120,13 @@ class BaseWhisperSTTService(SegmentedSTTService):
|
|||||||
model: str,
|
model: str,
|
||||||
api_key: Optional[str] = None,
|
api_key: Optional[str] = None,
|
||||||
base_url: Optional[str] = None,
|
base_url: Optional[str] = None,
|
||||||
|
language: Optional[Language] = Language.EN,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
self.set_model_name(model)
|
self.set_model_name(model)
|
||||||
self._client = self._create_client(api_key, base_url)
|
self._client = self._create_client(api_key, base_url)
|
||||||
|
self._language = self.language_to_service_language(language or Language.EN)
|
||||||
|
|
||||||
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)
|
||||||
@@ -57,6 +137,9 @@ class BaseWhisperSTTService(SegmentedSTTService):
|
|||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
def language_to_service_language(self, language: Language) -> Optional[str]:
|
||||||
|
return language_to_whisper_language(language)
|
||||||
|
|
||||||
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
async def run_stt(self, audio: bytes) -> AsyncGenerator[Frame, None]:
|
||||||
try:
|
try:
|
||||||
await self.start_processing_metrics()
|
await self.start_processing_metrics()
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from loguru import logger
|
|||||||
|
|
||||||
from pipecat.services.base_whisper import BaseWhisperSTTService, Transcription
|
from pipecat.services.base_whisper import BaseWhisperSTTService, Transcription
|
||||||
from pipecat.services.openai import OpenAILLMService
|
from pipecat.services.openai import OpenAILLMService
|
||||||
|
from pipecat.transcriptions.language import Language
|
||||||
|
|
||||||
|
|
||||||
class GroqLLMService(OpenAILLMService):
|
class GroqLLMService(OpenAILLMService):
|
||||||
@@ -52,8 +53,8 @@ class GroqSTTService(BaseWhisperSTTService):
|
|||||||
model: Whisper model to use. Defaults to "whisper-large-v3-turbo".
|
model: Whisper model to use. Defaults to "whisper-large-v3-turbo".
|
||||||
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. Defaults to English.
|
||||||
**kwargs: Additional arguments passed to BaseWhisperSTTService.
|
**kwargs: Additional arguments passed to BaseWhisperSTTService.
|
||||||
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -62,11 +63,18 @@ class GroqSTTService(BaseWhisperSTTService):
|
|||||||
model: str = "whisper-large-v3-turbo",
|
model: str = "whisper-large-v3-turbo",
|
||||||
api_key: Optional[str] = None,
|
api_key: Optional[str] = None,
|
||||||
base_url: str = "https://api.groq.com/openai/v1",
|
base_url: str = "https://api.groq.com/openai/v1",
|
||||||
|
language: Optional[Language] = Language.EN,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(model=model, api_key=api_key, base_url=base_url, **kwargs)
|
super().__init__(
|
||||||
|
model=model, api_key=api_key, base_url=base_url, language=language, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
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
|
||||||
return await self._client.audio.transcriptions.create(
|
return await self._client.audio.transcriptions.create(
|
||||||
file=("audio.wav", audio, "audio/wav"), model=self.model_name, response_format="json"
|
file=("audio.wav", audio, "audio/wav"),
|
||||||
|
model=self.model_name,
|
||||||
|
response_format="json",
|
||||||
|
language=self._language,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -54,6 +54,7 @@ from pipecat.services.ai_services import (
|
|||||||
TTSService,
|
TTSService,
|
||||||
)
|
)
|
||||||
from pipecat.services.base_whisper import BaseWhisperSTTService, Transcription
|
from pipecat.services.base_whisper import BaseWhisperSTTService, Transcription
|
||||||
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.utils.time import time_now_iso8601
|
from pipecat.utils.time import time_now_iso8601
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -406,6 +407,7 @@ class OpenAISTTService(BaseWhisperSTTService):
|
|||||||
model: Whisper model to use. Defaults to "whisper-1".
|
model: Whisper model to use. Defaults to "whisper-1".
|
||||||
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.
|
||||||
**kwargs: Additional arguments passed to BaseWhisperSTTService.
|
**kwargs: Additional arguments passed to BaseWhisperSTTService.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -415,13 +417,17 @@ class OpenAISTTService(BaseWhisperSTTService):
|
|||||||
model: str = "whisper-1",
|
model: str = "whisper-1",
|
||||||
api_key: Optional[str] = None,
|
api_key: Optional[str] = None,
|
||||||
base_url: Optional[str] = None,
|
base_url: Optional[str] = None,
|
||||||
|
language: Optional[Language] = Language.EN,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
super().__init__(model=model, api_key=api_key, base_url=base_url, **kwargs)
|
super().__init__(
|
||||||
|
model=model, api_key=api_key, base_url=base_url, language=language, **kwargs
|
||||||
|
)
|
||||||
|
|
||||||
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
|
||||||
return await self._client.audio.transcriptions.create(
|
return await self._client.audio.transcriptions.create(
|
||||||
file=("audio.wav", audio, "audio/wav"), model=self.model_name
|
file=("audio.wav", audio, "audio/wav"), model=self.model_name, language=self._language
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -54,6 +54,9 @@ class Language(StrEnum):
|
|||||||
AZ = "az"
|
AZ = "az"
|
||||||
AZ_AZ = "az-AZ"
|
AZ_AZ = "az-AZ"
|
||||||
|
|
||||||
|
# Belarusian
|
||||||
|
BE = "be"
|
||||||
|
|
||||||
# Bulgarian
|
# Bulgarian
|
||||||
BG = "bg"
|
BG = "bg"
|
||||||
BG_BG = "bg-BG"
|
BG_BG = "bg-BG"
|
||||||
@@ -264,6 +267,9 @@ class Language(StrEnum):
|
|||||||
MN = "mn"
|
MN = "mn"
|
||||||
MN_MN = "mn-MN"
|
MN_MN = "mn-MN"
|
||||||
|
|
||||||
|
# Maori
|
||||||
|
MI = "mi"
|
||||||
|
|
||||||
# Marathi
|
# Marathi
|
||||||
MR = "mr"
|
MR = "mr"
|
||||||
MR_IN = "mr-IN"
|
MR_IN = "mr-IN"
|
||||||
|
|||||||
Reference in New Issue
Block a user