services(together): fix together AI InputParams

This commit is contained in:
Aleix Conchillo Flaqué
2024-10-24 12:57:46 -07:00
parent 33553b71d4
commit d930a46e64
2 changed files with 5 additions and 39 deletions

View File

@@ -99,6 +99,9 @@ class BaseOpenAILLMService(LLMService):
) )
seed: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=0) seed: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=0)
temperature: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=2.0) temperature: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=2.0)
# Note: top_k is currently not supported by the OpenAI client library,
# so top_k is ignore right now.
top_k: Optional[int] = Field(default=None, ge=0)
top_p: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=1.0) top_p: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=1.0)
max_tokens: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=1) max_tokens: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=1)
max_completion_tokens: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=1) max_completion_tokens: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=1)

View File

@@ -4,11 +4,8 @@
# SPDX-License-Identifier: BSD 2-Clause License # SPDX-License-Identifier: BSD 2-Clause License
# #
from typing import Any, Dict, Optional
import httpx
from loguru import logger from loguru import logger
from pydantic import BaseModel, Field
from pipecat.services.openai import OpenAILLMService from pipecat.services.openai import OpenAILLMService
@@ -27,50 +24,16 @@ except ModuleNotFoundError as e:
class TogetherLLMService(OpenAILLMService): class TogetherLLMService(OpenAILLMService):
"""This class implements inference with Together's Llama 3.1 models""" """This class implements inference with Together's Llama 3.1 models"""
class InputParams(BaseModel):
frequency_penalty: Optional[float] = Field(default=None, ge=-2.0, le=2.0)
max_tokens: Optional[int] = Field(default=4096, ge=1)
presence_penalty: Optional[float] = Field(default=None, ge=-2.0, le=2.0)
temperature: Optional[float] = Field(default=None, ge=0.0, le=1.0)
# Note: top_k is currently not supported by the OpenAI client library,
# so top_k is ignore right now.
top_k: Optional[int] = Field(default=None, ge=0)
top_p: Optional[float] = Field(default=None, ge=0.0, le=1.0)
extra: Optional[Dict[str, Any]] = Field(default_factory=dict)
seed: Optional[int] = Field(default=None)
def __init__( def __init__(
self, self,
*, *,
api_key: str, api_key: str,
base_url: str = "https://api.together.xyz/v1", base_url: str = "https://api.together.xyz/v1",
model: str = "meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo", model: str = "meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo",
params: InputParams = InputParams(),
**kwargs, **kwargs,
): ):
super().__init__(api_key=api_key, base_url=base_url, model=model, params=params, **kwargs) super().__init__(api_key=api_key, base_url=base_url, model=model, **kwargs)
self.set_model_name(model)
self._settings = {
"max_tokens": params.max_tokens,
"frequency_penalty": params.frequency_penalty,
"presence_penalty": params.presence_penalty,
"seed": params.seed,
"temperature": params.temperature,
"top_p": params.top_p,
"extra": params.extra if isinstance(params.extra, dict) else {},
}
def can_generate_metrics(self) -> bool:
return True
def create_client(self, api_key=None, base_url=None, **kwargs): def create_client(self, api_key=None, base_url=None, **kwargs):
logger.debug(f"Creating Together.ai client with api {base_url}") logger.debug(f"Creating Together.ai client with api {base_url}")
return AsyncOpenAI( return super().create_client(api_key, base_url, **kwargs)
api_key=api_key,
base_url=base_url,
http_client=DefaultAsyncHttpxClient(
limits=httpx.Limits(
max_keepalive_connections=100, max_connections=1000, keepalive_expiry=None
)
),
)