mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(bedrock): freeze invoke params before request build
Parse Bedrock invoke params into provider-specific frozen dataclasses before constructing wire request bodies so builders cannot accept raw optional_params. Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
parent
ed082b1ef2
commit
73a4dcb9d9
2 changed files with 343 additions and 122 deletions
|
|
@ -1,12 +1,12 @@
|
|||
import copy
|
||||
import json
|
||||
import time
|
||||
from dataclasses import dataclass, replace
|
||||
from functools import partial
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Dict,
|
||||
FrozenSet,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
|
|
@ -36,14 +36,10 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
_get_httpx_client,
|
||||
)
|
||||
from litellm.types.llms.bedrock import (
|
||||
BedrockInvokeAI21InferenceParams,
|
||||
BedrockInvokeAI21Request,
|
||||
BedrockInvokeCohereChatRequest,
|
||||
BedrockInvokeCohereCompletionRequest,
|
||||
BedrockInvokeCohereInferenceParams,
|
||||
BedrockInvokeLlamaInferenceParams,
|
||||
BedrockInvokeLlamaRequest,
|
||||
BedrockInvokeMistralInferenceParams,
|
||||
BedrockInvokeMistralRequest,
|
||||
BedrockInvokeTitanInferenceParams,
|
||||
BedrockInvokeTitanRequest,
|
||||
|
|
@ -62,23 +58,63 @@ else:
|
|||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
|
||||
|
||||
class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
||||
BEDROCK_INVOKE_COHERE_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
|
||||
BedrockInvokeCohereInferenceParams.__annotations__.keys()
|
||||
)
|
||||
BEDROCK_INVOKE_AI21_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
|
||||
BedrockInvokeAI21InferenceParams.__annotations__.keys()
|
||||
)
|
||||
BEDROCK_INVOKE_MISTRAL_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
|
||||
BedrockInvokeMistralInferenceParams.__annotations__.keys()
|
||||
)
|
||||
BEDROCK_INVOKE_TITAN_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
|
||||
BedrockInvokeTitanInferenceParams.__annotations__.keys()
|
||||
)
|
||||
BEDROCK_INVOKE_LLAMA_ALLOWED_INFERENCE_FIELDS: FrozenSet[str] = frozenset(
|
||||
BedrockInvokeLlamaInferenceParams.__annotations__.keys()
|
||||
)
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BedrockInvokeCohereParams:
|
||||
max_tokens: Optional[int] = None
|
||||
temperature: Optional[float] = None
|
||||
return_likelihood: Optional[str] = None
|
||||
p: Optional[float] = None
|
||||
k: Optional[int] = None
|
||||
stop_sequences: Optional[List[str]] = None
|
||||
num_generations: Optional[int] = None
|
||||
frequency_penalty: Optional[float] = None
|
||||
presence_penalty: Optional[float] = None
|
||||
truncate: Optional[str] = None
|
||||
stream: Optional[bool] = None
|
||||
tools: Optional[List[Dict[str, object]]] = None
|
||||
tool_results: Optional[List[Dict[str, object]]] = None
|
||||
seed: Optional[int] = None
|
||||
force_single_step: Optional[bool] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BedrockInvokeAI21Params:
|
||||
maxTokens: Optional[int] = None
|
||||
temperature: Optional[float] = None
|
||||
topP: Optional[float] = None
|
||||
stopSequences: Optional[List[str]] = None
|
||||
frequencyPenalty: Optional[Dict[str, object]] = None
|
||||
frequencePenalty: Optional[Dict[str, object]] = None
|
||||
presencePenalty: Optional[Dict[str, object]] = None
|
||||
countPenalty: Optional[Dict[str, object]] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BedrockInvokeMistralParams:
|
||||
max_tokens: Optional[int] = None
|
||||
temperature: Optional[float] = None
|
||||
top_p: Optional[float] = None
|
||||
top_k: Optional[float] = None
|
||||
stop: Optional[List[str]] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BedrockInvokeTitanParams:
|
||||
maxTokenCount: Optional[int] = None
|
||||
stopSequences: Optional[List[str]] = None
|
||||
temperature: Optional[float] = None
|
||||
topP: Optional[int] = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BedrockInvokeLlamaParams:
|
||||
max_gen_len: Optional[int] = None
|
||||
temperature: Optional[float] = None
|
||||
top_p: Optional[float] = None
|
||||
topP: Optional[float] = None
|
||||
|
||||
|
||||
class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
||||
def __init__(self, **kwargs):
|
||||
BaseConfig.__init__(self, **kwargs)
|
||||
BaseAWSLLM.__init__(self, **kwargs)
|
||||
|
|
@ -174,94 +210,200 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
fake_stream=fake_stream,
|
||||
)
|
||||
|
||||
def _get_filtered_inference_params(
|
||||
self,
|
||||
optional_params: dict,
|
||||
config: dict,
|
||||
allowed_fields: FrozenSet[str],
|
||||
) -> Dict[str, object]:
|
||||
copied_params = copy.deepcopy(optional_params)
|
||||
inference_params = {
|
||||
k: v
|
||||
for k, v in copied_params.items()
|
||||
if k in allowed_fields and k not in self.aws_authentication_params
|
||||
}
|
||||
config_params = {
|
||||
k: v
|
||||
for k, v in config.items()
|
||||
if k in allowed_fields
|
||||
and k not in self.aws_authentication_params
|
||||
and k not in inference_params
|
||||
}
|
||||
return {**config_params, **inference_params}
|
||||
@staticmethod
|
||||
def _get_param_value(optional_params: dict, config: dict, key: str) -> object:
|
||||
if key in optional_params:
|
||||
return copy.deepcopy(optional_params[key])
|
||||
return copy.deepcopy(config.get(key))
|
||||
|
||||
def _get_cohere_inference_params(
|
||||
@staticmethod
|
||||
def _drop_none(data: Dict[str, object]) -> Dict[str, object]:
|
||||
return {key: value for key, value in data.items() if value is not None}
|
||||
|
||||
def _parse_cohere_params(
|
||||
self, optional_params: dict, config: dict
|
||||
) -> BedrockInvokeCohereInferenceParams:
|
||||
return cast(
|
||||
BedrockInvokeCohereInferenceParams,
|
||||
self._get_filtered_inference_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
allowed_fields=self.BEDROCK_INVOKE_COHERE_ALLOWED_INFERENCE_FIELDS,
|
||||
) -> BedrockInvokeCohereParams:
|
||||
return BedrockInvokeCohereParams(
|
||||
max_tokens=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "max_tokens"),
|
||||
),
|
||||
temperature=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "temperature"),
|
||||
),
|
||||
return_likelihood=cast(
|
||||
Optional[str],
|
||||
self._get_param_value(optional_params, config, "return_likelihood"),
|
||||
),
|
||||
p=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "p"),
|
||||
),
|
||||
k=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "k"),
|
||||
),
|
||||
stop_sequences=cast(
|
||||
Optional[List[str]],
|
||||
self._get_param_value(optional_params, config, "stop_sequences"),
|
||||
),
|
||||
num_generations=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "num_generations"),
|
||||
),
|
||||
frequency_penalty=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "frequency_penalty"),
|
||||
),
|
||||
presence_penalty=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "presence_penalty"),
|
||||
),
|
||||
truncate=cast(
|
||||
Optional[str],
|
||||
self._get_param_value(optional_params, config, "truncate"),
|
||||
),
|
||||
stream=cast(
|
||||
Optional[bool],
|
||||
self._get_param_value(optional_params, config, "stream"),
|
||||
),
|
||||
tools=cast(
|
||||
Optional[List[Dict[str, object]]],
|
||||
self._get_param_value(optional_params, config, "tools"),
|
||||
),
|
||||
tool_results=cast(
|
||||
Optional[List[Dict[str, object]]],
|
||||
self._get_param_value(optional_params, config, "tool_results"),
|
||||
),
|
||||
seed=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "seed"),
|
||||
),
|
||||
force_single_step=cast(
|
||||
Optional[bool],
|
||||
self._get_param_value(optional_params, config, "force_single_step"),
|
||||
),
|
||||
)
|
||||
|
||||
def _get_ai21_inference_params(
|
||||
def _parse_ai21_params(
|
||||
self, optional_params: dict, config: dict
|
||||
) -> BedrockInvokeAI21InferenceParams:
|
||||
return cast(
|
||||
BedrockInvokeAI21InferenceParams,
|
||||
self._get_filtered_inference_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
allowed_fields=self.BEDROCK_INVOKE_AI21_ALLOWED_INFERENCE_FIELDS,
|
||||
) -> BedrockInvokeAI21Params:
|
||||
return BedrockInvokeAI21Params(
|
||||
maxTokens=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "maxTokens"),
|
||||
),
|
||||
temperature=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "temperature"),
|
||||
),
|
||||
topP=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "topP"),
|
||||
),
|
||||
stopSequences=cast(
|
||||
Optional[List[str]],
|
||||
self._get_param_value(optional_params, config, "stopSequences"),
|
||||
),
|
||||
frequencyPenalty=cast(
|
||||
Optional[Dict[str, object]],
|
||||
self._get_param_value(optional_params, config, "frequencyPenalty"),
|
||||
),
|
||||
frequencePenalty=cast(
|
||||
Optional[Dict[str, object]],
|
||||
self._get_param_value(optional_params, config, "frequencePenalty"),
|
||||
),
|
||||
presencePenalty=cast(
|
||||
Optional[Dict[str, object]],
|
||||
self._get_param_value(optional_params, config, "presencePenalty"),
|
||||
),
|
||||
countPenalty=cast(
|
||||
Optional[Dict[str, object]],
|
||||
self._get_param_value(optional_params, config, "countPenalty"),
|
||||
),
|
||||
)
|
||||
|
||||
def _get_mistral_inference_params(
|
||||
def _parse_mistral_params(
|
||||
self, optional_params: dict, config: dict
|
||||
) -> BedrockInvokeMistralInferenceParams:
|
||||
return cast(
|
||||
BedrockInvokeMistralInferenceParams,
|
||||
self._get_filtered_inference_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
allowed_fields=self.BEDROCK_INVOKE_MISTRAL_ALLOWED_INFERENCE_FIELDS,
|
||||
) -> BedrockInvokeMistralParams:
|
||||
return BedrockInvokeMistralParams(
|
||||
max_tokens=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "max_tokens"),
|
||||
),
|
||||
temperature=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "temperature"),
|
||||
),
|
||||
top_p=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "top_p"),
|
||||
),
|
||||
top_k=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "top_k"),
|
||||
),
|
||||
stop=cast(
|
||||
Optional[List[str]],
|
||||
self._get_param_value(optional_params, config, "stop"),
|
||||
),
|
||||
)
|
||||
|
||||
def _get_titan_inference_params(
|
||||
def _parse_titan_params(
|
||||
self, optional_params: dict, config: dict
|
||||
) -> BedrockInvokeTitanInferenceParams:
|
||||
return cast(
|
||||
BedrockInvokeTitanInferenceParams,
|
||||
self._get_filtered_inference_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
allowed_fields=self.BEDROCK_INVOKE_TITAN_ALLOWED_INFERENCE_FIELDS,
|
||||
) -> BedrockInvokeTitanParams:
|
||||
return BedrockInvokeTitanParams(
|
||||
maxTokenCount=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "maxTokenCount"),
|
||||
),
|
||||
stopSequences=cast(
|
||||
Optional[List[str]],
|
||||
self._get_param_value(optional_params, config, "stopSequences"),
|
||||
),
|
||||
temperature=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "temperature"),
|
||||
),
|
||||
topP=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "topP"),
|
||||
),
|
||||
)
|
||||
|
||||
def _get_llama_inference_params(
|
||||
def _parse_llama_params(
|
||||
self, optional_params: dict, config: dict
|
||||
) -> BedrockInvokeLlamaInferenceParams:
|
||||
return cast(
|
||||
BedrockInvokeLlamaInferenceParams,
|
||||
self._get_filtered_inference_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
allowed_fields=self.BEDROCK_INVOKE_LLAMA_ALLOWED_INFERENCE_FIELDS,
|
||||
) -> BedrockInvokeLlamaParams:
|
||||
return BedrockInvokeLlamaParams(
|
||||
max_gen_len=cast(
|
||||
Optional[int],
|
||||
self._get_param_value(optional_params, config, "max_gen_len"),
|
||||
),
|
||||
temperature=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "temperature"),
|
||||
),
|
||||
top_p=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "top_p"),
|
||||
),
|
||||
topP=cast(
|
||||
Optional[float],
|
||||
self._get_param_value(optional_params, config, "topP"),
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_cohere_chat_request(
|
||||
prompt: str,
|
||||
inference_params: BedrockInvokeCohereInferenceParams,
|
||||
inference_params: BedrockInvokeCohereParams,
|
||||
chat_history: Optional[List[Dict[str, object]]],
|
||||
) -> BedrockInvokeCohereChatRequest:
|
||||
request: Dict[str, object] = {"message": prompt, **inference_params}
|
||||
request: Dict[str, object] = {
|
||||
"message": prompt,
|
||||
**AmazonInvokeConfig._cohere_params_to_wire_dict(inference_params),
|
||||
}
|
||||
if chat_history is not None:
|
||||
request["chat_history"] = chat_history
|
||||
return cast(BedrockInvokeCohereChatRequest, request)
|
||||
|
|
@ -269,48 +411,123 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
@staticmethod
|
||||
def _build_cohere_completion_request(
|
||||
prompt: str,
|
||||
inference_params: BedrockInvokeCohereInferenceParams,
|
||||
inference_params: BedrockInvokeCohereParams,
|
||||
) -> BedrockInvokeCohereCompletionRequest:
|
||||
return cast(
|
||||
BedrockInvokeCohereCompletionRequest,
|
||||
{"prompt": prompt, **inference_params},
|
||||
{
|
||||
"prompt": prompt,
|
||||
**AmazonInvokeConfig._cohere_params_to_wire_dict(inference_params),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_ai21_request(
|
||||
prompt: str,
|
||||
inference_params: BedrockInvokeAI21InferenceParams,
|
||||
inference_params: BedrockInvokeAI21Params,
|
||||
) -> BedrockInvokeAI21Request:
|
||||
return cast(BedrockInvokeAI21Request, {"prompt": prompt, **inference_params})
|
||||
return cast(
|
||||
BedrockInvokeAI21Request,
|
||||
{
|
||||
"prompt": prompt,
|
||||
**AmazonInvokeConfig._drop_none(
|
||||
{
|
||||
"maxTokens": inference_params.maxTokens,
|
||||
"temperature": inference_params.temperature,
|
||||
"topP": inference_params.topP,
|
||||
"stopSequences": inference_params.stopSequences,
|
||||
"frequencyPenalty": inference_params.frequencyPenalty,
|
||||
"frequencePenalty": inference_params.frequencePenalty,
|
||||
"presencePenalty": inference_params.presencePenalty,
|
||||
"countPenalty": inference_params.countPenalty,
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_mistral_request(
|
||||
prompt: str,
|
||||
inference_params: BedrockInvokeMistralInferenceParams,
|
||||
inference_params: BedrockInvokeMistralParams,
|
||||
) -> BedrockInvokeMistralRequest:
|
||||
return cast(
|
||||
BedrockInvokeMistralRequest,
|
||||
{"prompt": prompt, **inference_params},
|
||||
{
|
||||
"prompt": prompt,
|
||||
**AmazonInvokeConfig._drop_none(
|
||||
{
|
||||
"max_tokens": inference_params.max_tokens,
|
||||
"temperature": inference_params.temperature,
|
||||
"top_p": inference_params.top_p,
|
||||
"top_k": inference_params.top_k,
|
||||
"stop": inference_params.stop,
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_titan_request(
|
||||
prompt: str,
|
||||
inference_params: BedrockInvokeTitanInferenceParams,
|
||||
inference_params: BedrockInvokeTitanParams,
|
||||
) -> BedrockInvokeTitanRequest:
|
||||
return BedrockInvokeTitanRequest(
|
||||
inputText=prompt,
|
||||
textGenerationConfig=inference_params,
|
||||
textGenerationConfig=cast(
|
||||
BedrockInvokeTitanInferenceParams,
|
||||
AmazonInvokeConfig._drop_none(
|
||||
{
|
||||
"maxTokenCount": inference_params.maxTokenCount,
|
||||
"stopSequences": inference_params.stopSequences,
|
||||
"temperature": inference_params.temperature,
|
||||
"topP": inference_params.topP,
|
||||
}
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_llama_request(
|
||||
prompt: str,
|
||||
inference_params: BedrockInvokeLlamaInferenceParams,
|
||||
inference_params: BedrockInvokeLlamaParams,
|
||||
) -> BedrockInvokeLlamaRequest:
|
||||
return cast(
|
||||
BedrockInvokeLlamaRequest,
|
||||
{"prompt": prompt, **inference_params},
|
||||
{
|
||||
"prompt": prompt,
|
||||
**AmazonInvokeConfig._drop_none(
|
||||
{
|
||||
"max_gen_len": inference_params.max_gen_len,
|
||||
"temperature": inference_params.temperature,
|
||||
"top_p": inference_params.top_p,
|
||||
"topP": inference_params.topP,
|
||||
}
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _cohere_params_to_wire_dict(
|
||||
inference_params: BedrockInvokeCohereParams,
|
||||
) -> Dict[str, object]:
|
||||
return AmazonInvokeConfig._drop_none(
|
||||
{
|
||||
"max_tokens": inference_params.max_tokens,
|
||||
"temperature": inference_params.temperature,
|
||||
"return_likelihood": inference_params.return_likelihood,
|
||||
"p": inference_params.p,
|
||||
"k": inference_params.k,
|
||||
"stop_sequences": inference_params.stop_sequences,
|
||||
"num_generations": inference_params.num_generations,
|
||||
"frequency_penalty": inference_params.frequency_penalty,
|
||||
"presence_penalty": inference_params.presence_penalty,
|
||||
"truncate": inference_params.truncate,
|
||||
"stream": inference_params.stream,
|
||||
"tools": inference_params.tools,
|
||||
"tool_results": inference_params.tool_results,
|
||||
"seed": inference_params.seed,
|
||||
"force_single_step": inference_params.force_single_step,
|
||||
}
|
||||
)
|
||||
|
||||
def transform_request(
|
||||
|
|
@ -338,7 +555,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
if provider == "cohere":
|
||||
if model.startswith("cohere.command-r"):
|
||||
config = litellm.AmazonCohereChatConfig().get_config()
|
||||
cohere_inference_params = self._get_cohere_inference_params(
|
||||
cohere_inference_params = self._parse_cohere_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
)
|
||||
|
|
@ -352,12 +569,15 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
)
|
||||
else:
|
||||
config = litellm.AmazonCohereConfig.get_config()
|
||||
cohere_inference_params = self._get_cohere_inference_params(
|
||||
cohere_inference_params = self._parse_cohere_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
)
|
||||
if stream is True:
|
||||
cohere_inference_params["stream"] = True
|
||||
cohere_inference_params = replace(
|
||||
cohere_inference_params,
|
||||
stream=True,
|
||||
)
|
||||
return cast(
|
||||
dict,
|
||||
self._build_cohere_completion_request(
|
||||
|
|
@ -387,7 +607,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
)
|
||||
elif provider == "ai21":
|
||||
config = litellm.AmazonAI21Config.get_config()
|
||||
ai21_inference_params = self._get_ai21_inference_params(
|
||||
ai21_inference_params = self._parse_ai21_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
)
|
||||
|
|
@ -400,7 +620,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
)
|
||||
elif provider == "mistral":
|
||||
config = litellm.AmazonMistralConfig.get_config()
|
||||
mistral_inference_params = self._get_mistral_inference_params(
|
||||
mistral_inference_params = self._parse_mistral_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
)
|
||||
|
|
@ -413,7 +633,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
)
|
||||
elif provider == "amazon": # amazon titan
|
||||
config = litellm.AmazonTitanConfig.get_config()
|
||||
titan_inference_params = self._get_titan_inference_params(
|
||||
titan_inference_params = self._parse_titan_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
)
|
||||
|
|
@ -426,7 +646,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
)
|
||||
elif provider == "meta" or provider == "llama" or provider == "deepseek_r1":
|
||||
config = litellm.AmazonLlamaConfig.get_config()
|
||||
llama_inference_params = self._get_llama_inference_params(
|
||||
llama_inference_params = self._parse_llama_params(
|
||||
optional_params=optional_params,
|
||||
config=config,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -432,7 +432,8 @@ class BedrockInvokeNovaRequest(TypedDict, total=False):
|
|||
guardrailConfig: Optional[GuardrailConfigBlock]
|
||||
|
||||
|
||||
class BedrockInvokeCohereInferenceParams(TypedDict, total=False):
|
||||
class BedrockInvokeCohereCompletionRequest(TypedDict, total=False):
|
||||
prompt: Required[str]
|
||||
max_tokens: int
|
||||
temperature: float
|
||||
return_likelihood: str
|
||||
|
|
@ -450,18 +451,28 @@ class BedrockInvokeCohereInferenceParams(TypedDict, total=False):
|
|||
force_single_step: bool
|
||||
|
||||
|
||||
class BedrockInvokeCohereCompletionRequest(
|
||||
BedrockInvokeCohereInferenceParams, total=False
|
||||
):
|
||||
prompt: Required[str]
|
||||
|
||||
|
||||
class BedrockInvokeCohereChatRequest(BedrockInvokeCohereInferenceParams, total=False):
|
||||
class BedrockInvokeCohereChatRequest(TypedDict, total=False):
|
||||
message: Required[str]
|
||||
chat_history: List[Dict[str, object]]
|
||||
max_tokens: int
|
||||
temperature: float
|
||||
return_likelihood: str
|
||||
p: float
|
||||
k: int
|
||||
stop_sequences: List[str]
|
||||
num_generations: int
|
||||
frequency_penalty: float
|
||||
presence_penalty: float
|
||||
truncate: str
|
||||
stream: bool
|
||||
tools: List[Dict[str, object]]
|
||||
tool_results: List[Dict[str, object]]
|
||||
seed: int
|
||||
force_single_step: bool
|
||||
|
||||
|
||||
class BedrockInvokeAI21InferenceParams(TypedDict, total=False):
|
||||
class BedrockInvokeAI21Request(TypedDict, total=False):
|
||||
prompt: Required[str]
|
||||
maxTokens: int
|
||||
temperature: float
|
||||
topP: float
|
||||
|
|
@ -472,11 +483,8 @@ class BedrockInvokeAI21InferenceParams(TypedDict, total=False):
|
|||
countPenalty: Dict[str, object]
|
||||
|
||||
|
||||
class BedrockInvokeAI21Request(BedrockInvokeAI21InferenceParams, total=False):
|
||||
class BedrockInvokeMistralRequest(TypedDict, total=False):
|
||||
prompt: Required[str]
|
||||
|
||||
|
||||
class BedrockInvokeMistralInferenceParams(TypedDict, total=False):
|
||||
max_tokens: int
|
||||
temperature: float
|
||||
top_p: float
|
||||
|
|
@ -484,10 +492,6 @@ class BedrockInvokeMistralInferenceParams(TypedDict, total=False):
|
|||
stop: List[str]
|
||||
|
||||
|
||||
class BedrockInvokeMistralRequest(BedrockInvokeMistralInferenceParams, total=False):
|
||||
prompt: Required[str]
|
||||
|
||||
|
||||
class BedrockInvokeTitanInferenceParams(TypedDict, total=False):
|
||||
maxTokenCount: int
|
||||
stopSequences: List[str]
|
||||
|
|
@ -500,17 +504,14 @@ class BedrockInvokeTitanRequest(TypedDict):
|
|||
textGenerationConfig: BedrockInvokeTitanInferenceParams
|
||||
|
||||
|
||||
class BedrockInvokeLlamaInferenceParams(TypedDict, total=False):
|
||||
class BedrockInvokeLlamaRequest(TypedDict, total=False):
|
||||
prompt: Required[str]
|
||||
max_gen_len: int
|
||||
temperature: float
|
||||
top_p: float
|
||||
topP: float
|
||||
|
||||
|
||||
class BedrockInvokeLlamaRequest(BedrockInvokeLlamaInferenceParams, total=False):
|
||||
prompt: Required[str]
|
||||
|
||||
|
||||
class GenericStreamingChunk(TypedDict):
|
||||
text: Required[str]
|
||||
tool_use: Optional[ChatCompletionToolCallChunk]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue