services(azure): new AzureSTTService
This commit is contained in:
@@ -7,26 +7,30 @@
|
|||||||
import aiohttp
|
import aiohttp
|
||||||
import asyncio
|
import asyncio
|
||||||
import io
|
import io
|
||||||
|
import time
|
||||||
|
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from typing import AsyncGenerator
|
from typing import AsyncGenerator
|
||||||
|
|
||||||
from openai import AsyncAzureOpenAI
|
from pipecat.frames.frames import AudioRawFrame, CancelFrame, EndFrame, ErrorFrame, Frame, StartFrame, SystemFrame, TranscriptionFrame, URLImageRawFrame
|
||||||
|
from pipecat.processors.frame_processor import FrameDirection
|
||||||
from pipecat.frames.frames import AudioRawFrame, ErrorFrame, Frame, URLImageRawFrame
|
from pipecat.services.ai_services import AIService, TTSService, ImageGenService
|
||||||
from pipecat.services.ai_services import TTSService, ImageGenService
|
|
||||||
from pipecat.services.openai import BaseOpenAILLMService
|
from pipecat.services.openai import BaseOpenAILLMService
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
# See .env.example for Azure configuration needed
|
# See .env.example for Azure configuration needed
|
||||||
try:
|
try:
|
||||||
|
from openai import AsyncAzureOpenAI
|
||||||
from azure.cognitiveservices.speech import (
|
from azure.cognitiveservices.speech import (
|
||||||
SpeechSynthesizer,
|
|
||||||
SpeechConfig,
|
SpeechConfig,
|
||||||
|
SpeechRecognizer,
|
||||||
|
SpeechSynthesizer,
|
||||||
ResultReason,
|
ResultReason,
|
||||||
CancellationReason,
|
CancellationReason,
|
||||||
)
|
)
|
||||||
|
from azure.cognitiveservices.speech.audio import AudioStreamFormat, PushAudioInputStream
|
||||||
|
from azure.cognitiveservices.speech.dialog import AudioConfig
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
logger.error(f"Exception: {e}")
|
logger.error(f"Exception: {e}")
|
||||||
logger.error(
|
logger.error(
|
||||||
@@ -34,14 +38,35 @@ except ModuleNotFoundError as e:
|
|||||||
raise Exception(f"Missing module: {e}")
|
raise Exception(f"Missing module: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
class AzureLLMService(BaseOpenAILLMService):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
api_key: str,
|
||||||
|
endpoint: str,
|
||||||
|
model: str,
|
||||||
|
api_version: str = "2023-12-01-preview"):
|
||||||
|
# Initialize variables before calling parent __init__() because that
|
||||||
|
# will call create_client() and we need those values there.
|
||||||
|
self._endpoint = endpoint
|
||||||
|
self._api_version = api_version
|
||||||
|
super().__init__(api_key=api_key, model=model)
|
||||||
|
|
||||||
|
def create_client(self, api_key=None, base_url=None, **kwargs):
|
||||||
|
return AsyncAzureOpenAI(
|
||||||
|
api_key=api_key,
|
||||||
|
azure_endpoint=self._endpoint,
|
||||||
|
api_version=self._api_version,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class AzureTTSService(TTSService):
|
class AzureTTSService(TTSService):
|
||||||
def __init__(self, *, api_key: str, region: str, voice="en-US-SaraNeural", **kwargs):
|
def __init__(self, *, api_key: str, region: str, voice="en-US-SaraNeural", **kwargs):
|
||||||
super().__init__(**kwargs)
|
super().__init__(**kwargs)
|
||||||
|
|
||||||
self.speech_config = SpeechConfig(subscription=api_key, region=region)
|
speech_config = SpeechConfig(subscription=api_key, region=region)
|
||||||
self.speech_synthesizer = SpeechSynthesizer(
|
self._speech_synthesizer = SpeechSynthesizer(speech_config=speech_config, audio_config=None)
|
||||||
speech_config=self.speech_config, audio_config=None
|
|
||||||
)
|
|
||||||
self._voice = voice
|
self._voice = voice
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
@@ -62,7 +87,7 @@ class AzureTTSService(TTSService):
|
|||||||
f"{text}"
|
f"{text}"
|
||||||
"</prosody></mstts:express-as></voice></speak> ")
|
"</prosody></mstts:express-as></voice></speak> ")
|
||||||
|
|
||||||
result = await asyncio.to_thread(self.speech_synthesizer.speak_ssml, (ssml))
|
result = await asyncio.to_thread(self._speech_synthesizer.speak_ssml, (ssml))
|
||||||
|
|
||||||
if result.reason == ResultReason.SynthesizingAudioCompleted:
|
if result.reason == ResultReason.SynthesizingAudioCompleted:
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
@@ -75,26 +100,73 @@ class AzureTTSService(TTSService):
|
|||||||
logger.error(f"{self} error: {cancellation_details.error_details}")
|
logger.error(f"{self} error: {cancellation_details.error_details}")
|
||||||
|
|
||||||
|
|
||||||
class AzureLLMService(BaseOpenAILLMService):
|
class AzureSTTService(AIService):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
endpoint: str,
|
region: str,
|
||||||
model: str,
|
language="en-US",
|
||||||
api_version: str = "2023-12-01-preview"):
|
sample_rate=16000,
|
||||||
# Initialize variables before calling parent __init__() because that
|
channels=1,
|
||||||
# will call create_client() and we need those values there.
|
**kwargs):
|
||||||
self._endpoint = endpoint
|
super().__init__(**kwargs)
|
||||||
self._api_version = api_version
|
|
||||||
super().__init__(api_key=api_key, model=model)
|
|
||||||
|
|
||||||
def create_client(self, api_key=None, base_url=None):
|
speech_config = SpeechConfig(subscription=api_key, region=region)
|
||||||
return AsyncAzureOpenAI(
|
speech_config.speech_recognition_language = language
|
||||||
api_key=api_key,
|
|
||||||
azure_endpoint=self._endpoint,
|
stream_format = AudioStreamFormat(samples_per_second=sample_rate, channels=channels)
|
||||||
api_version=self._api_version,
|
self._audio_stream = PushAudioInputStream(stream_format)
|
||||||
)
|
|
||||||
|
audio_config = AudioConfig(stream=self._audio_stream)
|
||||||
|
self._speech_recognizer = SpeechRecognizer(
|
||||||
|
speech_config=speech_config, audio_config=audio_config)
|
||||||
|
self._speech_recognizer.recognized.connect(self._on_handle_recognized)
|
||||||
|
|
||||||
|
self._create_push_task()
|
||||||
|
|
||||||
|
async def process_frame(self, frame: Frame, direction: FrameDirection):
|
||||||
|
await super().process_frame(frame, direction)
|
||||||
|
|
||||||
|
if isinstance(frame, SystemFrame):
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
elif isinstance(frame, AudioRawFrame):
|
||||||
|
self._audio_stream.write(frame.audio)
|
||||||
|
else:
|
||||||
|
await self._push_queue.put((frame, direction))
|
||||||
|
|
||||||
|
async def start(self, frame: StartFrame):
|
||||||
|
self._speech_recognizer.start_continuous_recognition_async()
|
||||||
|
|
||||||
|
async def stop(self, frame: EndFrame):
|
||||||
|
self._speech_recognizer.stop_continuous_recognition_async()
|
||||||
|
await self._push_queue.put((frame, FrameDirection.DOWNSTREAM))
|
||||||
|
await self._push_frame_task
|
||||||
|
|
||||||
|
async def cancel(self, frame: CancelFrame):
|
||||||
|
self._speech_recognizer.stop_continuous_recognition_async()
|
||||||
|
self._push_frame_task.cancel()
|
||||||
|
|
||||||
|
def _create_push_task(self):
|
||||||
|
self._push_frame_task = self.get_event_loop().create_task(self._push_frame_task_handler())
|
||||||
|
self._push_queue = asyncio.Queue()
|
||||||
|
|
||||||
|
async def _push_frame_task_handler(self):
|
||||||
|
running = True
|
||||||
|
while running:
|
||||||
|
try:
|
||||||
|
(frame, direction) = await self._push_queue.get()
|
||||||
|
await self.push_frame(frame, direction)
|
||||||
|
running = not isinstance(frame, EndFrame)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
break
|
||||||
|
|
||||||
|
def _on_handle_recognized(self, event):
|
||||||
|
if event.result.reason == ResultReason.RecognizedSpeech and len(event.result.text) > 0:
|
||||||
|
direction = FrameDirection.DOWNSTREAM
|
||||||
|
frame = TranscriptionFrame(event.result.text, "", int(time.time_ns() / 1000000))
|
||||||
|
asyncio.run_coroutine_threadsafe(
|
||||||
|
self._push_queue.put((frame, direction)), self.get_event_loop())
|
||||||
|
|
||||||
|
|
||||||
class AzureImageGenServiceREST(ImageGenService):
|
class AzureImageGenServiceREST(ImageGenService):
|
||||||
|
|||||||
Reference in New Issue
Block a user