Merge pull request #174 from pipecat-ai/fix-azure-llm-service
services(azure): fix AzureLLMService
This commit is contained in:
10
CHANGELOG.md
10
CHANGELOG.md
@@ -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
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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__()
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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]:
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user