Merge pull request #174 from pipecat-ai/fix-azure-llm-service

services(azure): fix AzureLLMService
This commit is contained in:
Aleix Conchillo Flaqué
2024-05-25 00:27:51 +08:00
committed by GitHub
10 changed files with 38 additions and 46 deletions

View File

@@ -5,6 +5,16 @@ All notable changes to **pipecat** will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/), The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html). and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
## [Unreleased]
### Changed
- GoogleLLMService `api_key` argument is now mandatory.
### Fixed
- Fixed AzureLLMService.
## [0.0.23] - 2024-05-23 ## [0.0.23] - 2024-05-23
### Fixed ### Fixed

View File

@@ -62,19 +62,15 @@ async def main(room_url: str, token):
) )
) )
tts = ElevenLabsTTSService(
aiohttp_session=session,
api_key=os.getenv("ELEVENLABS_API_KEY"),
voice_id=os.getenv("ELEVENLABS_VOICE_ID"),
)
user_response = UserResponseAggregator() user_response = UserResponseAggregator()
image_requester = UserImageRequester() image_requester = UserImageRequester()
vision_aggregator = VisionImageFrameAggregator() vision_aggregator = VisionImageFrameAggregator()
google = GoogleLLMService(model="gemini-1.5-flash-latest") google = GoogleLLMService(
model="gemini-1.5-flash-latest",
api_key=os.getenv("GOOGLE_API_KEY"))
tts = ElevenLabsTTSService( tts = ElevenLabsTTSService(
aiohttp_session=session, aiohttp_session=session,

View File

@@ -61,12 +61,6 @@ async def main(room_url: str, token):
) )
) )
tts = ElevenLabsTTSService(
aiohttp_session=session,
api_key=os.getenv("ELEVENLABS_API_KEY"),
voice_id=os.getenv("ELEVENLABS_VOICE_ID"),
)
user_response = UserResponseAggregator() user_response = UserResponseAggregator()
image_requester = UserImageRequester() image_requester = UserImageRequester()

View File

@@ -61,12 +61,6 @@ async def main(room_url: str, token):
) )
) )
tts = ElevenLabsTTSService(
aiohttp_session=session,
api_key=os.getenv("ELEVENLABS_API_KEY"),
voice_id=os.getenv("ELEVENLABS_VOICE_ID"),
)
user_response = UserResponseAggregator() user_response = UserResponseAggregator()
image_requester = UserImageRequester() image_requester = UserImageRequester()

View File

@@ -44,9 +44,9 @@ class AnthropicLLMService(LLMService):
def __init__( def __init__(
self, self,
api_key, api_key: str,
model="claude-3-opus-20240229", model: str = "claude-3-opus-20240229",
max_tokens=1024): max_tokens: int = 1024):
super().__init__() super().__init__()
self._client = AsyncAnthropic(api_key=api_key) self._client = AsyncAnthropic(api_key=api_key)
self._model = model self._model = model

View File

@@ -11,6 +11,7 @@ import io
from PIL import Image from PIL import Image
from typing import AsyncGenerator from typing import AsyncGenerator
from numpy import str_
from openai import AsyncAzureOpenAI from openai import AsyncAzureOpenAI
from pipecat.frames.frames import AudioRawFrame, ErrorFrame, Frame, URLImageRawFrame from pipecat.frames.frames import AudioRawFrame, ErrorFrame, Frame, URLImageRawFrame
@@ -73,17 +74,18 @@ class AzureLLMService(BaseOpenAILLMService):
def __init__( def __init__(
self, self,
*, *,
api_key, api_key: str,
endpoint, endpoint: str,
api_version="2023-12-01-preview", model: str,
model): api_version: str = "2023-12-01-preview"):
super().__init__(api_key=api_key, model=model) # Initialize variables before calling parent __init__() because that
# will call create_client() and we need those values there.
self._endpoint = endpoint self._endpoint = endpoint
self._api_version = api_version self._api_version = api_version
self._model: str = model super().__init__(api_key=api_key, model=model)
def create_client(self, api_key=None, base_url=None): def create_client(self, api_key=None, base_url=None):
self._client = AsyncAzureOpenAI( return AsyncAzureOpenAI(
api_key=api_key, api_key=api_key,
azure_endpoint=self._endpoint, azure_endpoint=self._endpoint,
api_version=self._api_version, api_version=self._api_version,
@@ -95,12 +97,12 @@ class AzureImageGenServiceREST(ImageGenService):
def __init__( def __init__(
self, self,
*, *,
api_version="2023-06-01-preview",
image_size: str,
aiohttp_session: aiohttp.ClientSession, aiohttp_session: aiohttp.ClientSession,
api_key, image_size: str,
endpoint, api_key: str,
model, endpoint: str,
model: str,
api_version="2023-06-01-preview",
): ):
super().__init__() super().__init__()

View File

@@ -19,6 +19,6 @@ except ModuleNotFoundError as e:
class FireworksLLMService(BaseOpenAILLMService): class FireworksLLMService(BaseOpenAILLMService):
def __init__(self, def __init__(self,
model="accounts/fireworks/models/firefunction-v1", model: str = "accounts/fireworks/models/firefunction-v1",
base_url="https://api.fireworks.ai/inference/v1"): base_url: str = "https://api.fireworks.ai/inference/v1"):
super().__init__(model, base_url) super().__init__(model, base_url)

View File

@@ -40,14 +40,10 @@ class GoogleLLMService(LLMService):
franca for all LLM services, so that it is easy to switch between different LLMs. franca for all LLM services, so that it is easy to switch between different LLMs.
""" """
def __init__(self, model="gemini-1.5-flash-latest", api_key=None, **kwargs): def __init__(self, api_key: str, model: str = "gemini-1.5-flash-latest", **kwargs):
super().__init__(**kwargs) super().__init__(**kwargs)
self.model = model gai.configure(api_key=api_key)
gai.configure(api_key=api_key or os.environ["GOOGLE_API_KEY"]) self._client = gai.GenerativeModel(model)
self.create_client()
def create_client(self):
self._client = gai.GenerativeModel(self.model)
def _get_messages_from_openai_context( def _get_messages_from_openai_context(
self, context: OpenAILLMContext) -> List[glm.Content]: self, context: OpenAILLMContext) -> List[glm.Content]:

View File

@@ -9,5 +9,5 @@ from pipecat.services.openai import BaseOpenAILLMService
class OLLamaLLMService(BaseOpenAILLMService): class OLLamaLLMService(BaseOpenAILLMService):
def __init__(self, model="llama2", base_url="http://localhost:11434/v1"): def __init__(self, model: str = "llama2", base_url: str = "http://localhost:11434/v1"):
super().__init__(model=model, base_url=base_url, api_key="ollama") super().__init__(model=model, base_url=base_url, api_key="ollama")

View File

@@ -58,10 +58,10 @@ class BaseOpenAILLMService(LLMService):
def __init__(self, model: str, api_key=None, base_url=None): def __init__(self, model: str, api_key=None, base_url=None):
super().__init__() super().__init__()
self._model: str = model self._model: str = model
self.create_client(api_key=api_key, base_url=base_url) self._client = self.create_client(api_key=api_key, base_url=base_url)
def create_client(self, api_key=None, base_url=None): def create_client(self, api_key=None, base_url=None):
self._client = AsyncOpenAI(api_key=api_key, base_url=base_url) return AsyncOpenAI(api_key=api_key, base_url=base_url)
async def _stream_chat_completions( async def _stream_chat_completions(
self, context: OpenAILLMContext self, context: OpenAILLMContext