Add extra input param to LLMs
This commit is contained in:
@@ -57,10 +57,12 @@ async def main():
|
|||||||
model=os.getenv("TOGETHER_MODEL"),
|
model=os.getenv("TOGETHER_MODEL"),
|
||||||
params=TogetherLLMService.InputParams(
|
params=TogetherLLMService.InputParams(
|
||||||
temperature=1.0,
|
temperature=1.0,
|
||||||
frequency_penalty=2.0,
|
|
||||||
presence_penalty=0.0,
|
|
||||||
top_p=0.9,
|
top_p=0.9,
|
||||||
top_k=40
|
top_k=40,
|
||||||
|
extra={
|
||||||
|
"frequency_penalty": 2.0,
|
||||||
|
"presence_penalty": 0.0,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ import base64
|
|||||||
import json
|
import json
|
||||||
import io
|
import io
|
||||||
import copy
|
import copy
|
||||||
from typing import List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
from asyncio import CancelledError
|
from asyncio import CancelledError
|
||||||
@@ -81,6 +81,7 @@ class AnthropicLLMService(LLMService):
|
|||||||
temperature: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=1.0)
|
temperature: Optional[float] = Field(default_factory=lambda: NOT_GIVEN, ge=0.0, le=1.0)
|
||||||
top_k: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, ge=0)
|
top_k: Optional[int] = Field(default_factory=lambda: NOT_GIVEN, 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)
|
||||||
|
extra: Optional[Dict[str, Any]] = Field(default_factory=dict)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -97,6 +98,7 @@ class AnthropicLLMService(LLMService):
|
|||||||
self._temperature = params.temperature
|
self._temperature = params.temperature
|
||||||
self._top_k = params.top_k
|
self._top_k = params.top_k
|
||||||
self._top_p = params.top_p
|
self._top_p = params.top_p
|
||||||
|
self._extra = params.extra if isinstance(params.extra, dict) else {}
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -134,6 +136,10 @@ class AnthropicLLMService(LLMService):
|
|||||||
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
||||||
self._top_p = top_p
|
self._top_p = top_p
|
||||||
|
|
||||||
|
async def set_extra(self, extra: Dict[str, Any]):
|
||||||
|
logger.debug(f"Switching LLM extra to: [{extra}]")
|
||||||
|
self._extra = extra
|
||||||
|
|
||||||
async def _process_context(self, context: OpenAILLMContext):
|
async def _process_context(self, context: OpenAILLMContext):
|
||||||
# Usage tracking. We track the usage reported by Anthropic in prompt_tokens and
|
# Usage tracking. We track the usage reported by Anthropic in prompt_tokens and
|
||||||
# completion_tokens. We also estimate the completion tokens from output text
|
# completion_tokens. We also estimate the completion tokens from output text
|
||||||
@@ -163,16 +169,21 @@ class AnthropicLLMService(LLMService):
|
|||||||
|
|
||||||
await self.start_ttfb_metrics()
|
await self.start_ttfb_metrics()
|
||||||
|
|
||||||
response = await api_call(
|
params = {
|
||||||
tools=context.tools or [],
|
"tools": context.tools or [],
|
||||||
system=context.system,
|
"system": context.system,
|
||||||
messages=messages,
|
"messages": messages,
|
||||||
model=self.model_name,
|
"model": self.model_name,
|
||||||
max_tokens=self._max_tokens,
|
"max_tokens": self._max_tokens,
|
||||||
stream=True,
|
"stream": True,
|
||||||
temperature=self._temperature,
|
"temperature": self._temperature,
|
||||||
top_k=self._top_k,
|
"top_k": self._top_k,
|
||||||
top_p=self._top_p)
|
"top_p": self._top_p
|
||||||
|
}
|
||||||
|
|
||||||
|
params.update(self._extra)
|
||||||
|
|
||||||
|
response = await api_call(**params)
|
||||||
|
|
||||||
await self.stop_ttfb_metrics()
|
await self.stop_ttfb_metrics()
|
||||||
|
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ import json
|
|||||||
import httpx
|
import httpx
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
|
||||||
from typing import AsyncGenerator, Dict, List, Literal, Optional
|
from typing import Any, AsyncGenerator, Dict, List, Literal, Optional
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@@ -90,6 +90,7 @@ 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)
|
||||||
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)
|
||||||
|
extra: Optional[Dict[str, Any]] = Field(default_factory=dict)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -107,6 +108,7 @@ class BaseOpenAILLMService(LLMService):
|
|||||||
self._seed = params.seed
|
self._seed = params.seed
|
||||||
self._temperature = params.temperature
|
self._temperature = params.temperature
|
||||||
self._top_p = params.top_p
|
self._top_p = params.top_p
|
||||||
|
self._extra = params.extra if isinstance(params.extra, dict) else {}
|
||||||
|
|
||||||
def create_client(self, api_key=None, base_url=None, **kwargs):
|
def create_client(self, api_key=None, base_url=None, **kwargs):
|
||||||
return AsyncOpenAI(
|
return AsyncOpenAI(
|
||||||
@@ -141,23 +143,32 @@ class BaseOpenAILLMService(LLMService):
|
|||||||
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
||||||
self._top_p = top_p
|
self._top_p = top_p
|
||||||
|
|
||||||
|
async def set_extra(self, extra: Dict[str, Any]):
|
||||||
|
logger.debug(f"Switching LLM extra to: [{extra}]")
|
||||||
|
self._extra = extra
|
||||||
|
|
||||||
async def get_chat_completions(
|
async def get_chat_completions(
|
||||||
self,
|
self,
|
||||||
context: OpenAILLMContext,
|
context: OpenAILLMContext,
|
||||||
messages: List[ChatCompletionMessageParam]) -> AsyncStream[ChatCompletionChunk]:
|
messages: List[ChatCompletionMessageParam]) -> AsyncStream[ChatCompletionChunk]:
|
||||||
chunks = await self._client.chat.completions.create(
|
|
||||||
model=self.model_name,
|
params = {
|
||||||
stream=True,
|
"model": self.model_name,
|
||||||
messages=messages,
|
"stream": True,
|
||||||
tools=context.tools,
|
"messages": messages,
|
||||||
tool_choice=context.tool_choice,
|
"tools": context.tools,
|
||||||
stream_options={"include_usage": True},
|
"tool_choice": context.tool_choice,
|
||||||
frequency_penalty=self._frequency_penalty,
|
"stream_options": {"include_usage": True},
|
||||||
presence_penalty=self._presence_penalty,
|
"frequency_penalty": self._frequency_penalty,
|
||||||
seed=self._seed,
|
"presence_penalty": self._presence_penalty,
|
||||||
temperature=self._temperature,
|
"seed": self._seed,
|
||||||
top_p=self._top_p
|
"temperature": self._temperature,
|
||||||
)
|
"top_p": self._top_p,
|
||||||
|
}
|
||||||
|
|
||||||
|
params.update(self._extra)
|
||||||
|
|
||||||
|
chunks = await self._client.chat.completions.create(**params)
|
||||||
return chunks
|
return chunks
|
||||||
|
|
||||||
async def _stream_chat_completions(
|
async def _stream_chat_completions(
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ import re
|
|||||||
import uuid
|
import uuid
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
from typing import List
|
from typing import Any, Dict, List, Optional
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from asyncio import CancelledError
|
from asyncio import CancelledError
|
||||||
|
|
||||||
@@ -64,6 +64,7 @@ class TogetherLLMService(LLMService):
|
|||||||
temperature: Optional[float] = Field(default=None, ge=0.0, le=1.0)
|
temperature: Optional[float] = Field(default=None, ge=0.0, le=1.0)
|
||||||
top_k: Optional[int] = Field(default=None, ge=0)
|
top_k: Optional[int] = Field(default=None, ge=0)
|
||||||
top_p: Optional[float] = Field(default=None, ge=0.0, le=1.0)
|
top_p: Optional[float] = Field(default=None, ge=0.0, le=1.0)
|
||||||
|
extra: Optional[Dict[str, Any]] = Field(default_factory=dict)
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -81,6 +82,7 @@ class TogetherLLMService(LLMService):
|
|||||||
self._temperature = params.temperature
|
self._temperature = params.temperature
|
||||||
self._top_k = params.top_k
|
self._top_k = params.top_k
|
||||||
self._top_p = params.top_p
|
self._top_p = params.top_p
|
||||||
|
self._extra = params.extra if isinstance(params.extra, dict) else {}
|
||||||
|
|
||||||
def can_generate_metrics(self) -> bool:
|
def can_generate_metrics(self) -> bool:
|
||||||
return True
|
return True
|
||||||
@@ -118,6 +120,10 @@ class TogetherLLMService(LLMService):
|
|||||||
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
logger.debug(f"Switching LLM top_p to: [{top_p}]")
|
||||||
self._top_p = top_p
|
self._top_p = top_p
|
||||||
|
|
||||||
|
async def set_extra(self, extra: Dict[str, Any]):
|
||||||
|
logger.debug(f"Switching LLM extra to: [{extra}]")
|
||||||
|
self._extra = extra
|
||||||
|
|
||||||
async def _process_context(self, context: OpenAILLMContext):
|
async def _process_context(self, context: OpenAILLMContext):
|
||||||
try:
|
try:
|
||||||
await self.push_frame(LLMFullResponseStartFrame())
|
await self.push_frame(LLMFullResponseStartFrame())
|
||||||
@@ -127,17 +133,21 @@ class TogetherLLMService(LLMService):
|
|||||||
|
|
||||||
await self.start_ttfb_metrics()
|
await self.start_ttfb_metrics()
|
||||||
|
|
||||||
stream = await self._client.chat.completions.create(
|
params = {
|
||||||
messages=context.messages,
|
"messages": context.messages,
|
||||||
model=self.model_name,
|
"model": self.model_name,
|
||||||
max_tokens=self._max_tokens,
|
"max_tokens": self._max_tokens,
|
||||||
stream=True,
|
"stream": True,
|
||||||
frequency_penalty=self._frequency_penalty,
|
"frequency_penalty": self._frequency_penalty,
|
||||||
presence_penalty=self._presence_penalty,
|
"presence_penalty": self._presence_penalty,
|
||||||
temperature=self._temperature,
|
"temperature": self._temperature,
|
||||||
top_k=self._top_k,
|
"top_k": self._top_k,
|
||||||
top_p=self._top_p
|
"top_p": self._top_p
|
||||||
)
|
}
|
||||||
|
|
||||||
|
params.update(self._extra)
|
||||||
|
|
||||||
|
stream = await self._client.chat.completions.create(**params)
|
||||||
|
|
||||||
# Function calling
|
# Function calling
|
||||||
got_first_chunk = False
|
got_first_chunk = False
|
||||||
|
|||||||
Reference in New Issue
Block a user