fix(utils.py): fix vertex ai optional param handling

don't pass max retries to unsupported route

Fixes https://github.com/BerriAI/litellm/issues/8254
This commit is contained in:
Krrish Dholakia 2025-02-11 18:07:55 -08:00
parent 5e58ae0347
commit 04d07b090e
3 changed files with 55 additions and 33 deletions

View file

@ -2,9 +2,10 @@ import types
from typing import Optional
import litellm
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
class VertexAILlama3Config:
class VertexAILlama3Config(OpenAIGPTConfig):
"""
Reference:https://cloud.google.com/vertex-ai/generative-ai/docs/partner-models/llama#streaming
@ -46,8 +47,13 @@ class VertexAILlama3Config:
and v is not None
}
def get_supported_openai_params(self):
return litellm.OpenAIConfig().get_supported_openai_params(model="gpt-3.5-turbo")
def get_supported_openai_params(self, model: str):
supported_params = super().get_supported_openai_params(model=model)
try:
supported_params.remove("max_retries")
except KeyError:
pass
return supported_params
def map_openai_params(
self,
@ -60,7 +66,7 @@ class VertexAILlama3Config:
non_default_params["max_tokens"] = non_default_params.pop(
"max_completion_tokens"
)
return litellm.OpenAIConfig().map_openai_params(
return super().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,

View file

@ -3166,51 +3166,56 @@ def get_optional_params( # noqa: PLR0915
else False
),
)
elif custom_llm_provider == "vertex_ai" and model in litellm.vertex_llama3_models:
optional_params = litellm.VertexAILlama3Config().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=(
drop_params
if drop_params is not None and isinstance(drop_params, bool)
else False
),
)
elif custom_llm_provider == "vertex_ai" and model in litellm.vertex_mistral_models:
if "codestral" in model:
optional_params = litellm.CodestralTextCompletionConfig().map_openai_params(
model=model,
elif custom_llm_provider == "vertex_ai":
if model in litellm.vertex_mistral_models:
if "codestral" in model:
optional_params = (
litellm.CodestralTextCompletionConfig().map_openai_params(
model=model,
non_default_params=non_default_params,
optional_params=optional_params,
drop_params=(
drop_params
if drop_params is not None and isinstance(drop_params, bool)
else False
),
)
)
else:
optional_params = litellm.MistralConfig().map_openai_params(
model=model,
non_default_params=non_default_params,
optional_params=optional_params,
drop_params=(
drop_params
if drop_params is not None and isinstance(drop_params, bool)
else False
),
)
elif model in litellm.vertex_ai_ai21_models:
optional_params = litellm.VertexAIAi21Config().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=(
drop_params
if drop_params is not None and isinstance(drop_params, bool)
else False
),
)
else:
optional_params = litellm.MistralConfig().map_openai_params(
model=model,
else: # use generic openai-like param mapping
optional_params = litellm.VertexAILlama3Config().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=(
drop_params
if drop_params is not None and isinstance(drop_params, bool)
else False
),
)
elif custom_llm_provider == "vertex_ai" and model in litellm.vertex_ai_ai21_models:
optional_params = litellm.VertexAIAi21Config().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=(
drop_params
if drop_params is not None and isinstance(drop_params, bool)
else False
),
)
elif custom_llm_provider == "sagemaker":
# temperature, top_p, n, stream, stop, max_tokens, n, presence_penalty default to None
optional_params = litellm.SagemakerConfig().map_openai_params(

View file

@ -1067,3 +1067,14 @@ def test_gemini_frequency_penalty():
model="gemini-1.5-flash", custom_llm_provider="gemini", frequency_penalty=0.5
)
assert optional_params["frequency_penalty"] == 0.5
def test_vertex_ai_ft_llama():
optional_params = get_optional_params(
model="1984786713414729728",
custom_llm_provider="vertex_ai",
frequency_penalty=0.5,
max_retries=10,
)
assert optional_params["frequency_penalty"] == 0.5
assert "max_retries" not in optional_params