From 04d07b090e21fbc1d038dee03b51b7b41d926266 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Tue, 11 Feb 2025 18:07:55 -0800 Subject: [PATCH] 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 --- .../llama3/transformation.py | 14 +++-- litellm/utils.py | 63 ++++++++++--------- tests/llm_translation/test_optional_params.py | 11 ++++ 3 files changed, 55 insertions(+), 33 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py index 0278d198062..4a8786aed1c 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py @@ -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, diff --git a/litellm/utils.py b/litellm/utils.py index dbefc90dc61..34a5dc596cd 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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( diff --git a/tests/llm_translation/test_optional_params.py b/tests/llm_translation/test_optional_params.py index 75fe7aa5b19..e7f2f8ac280 100644 --- a/tests/llm_translation/test_optional_params.py +++ b/tests/llm_translation/test_optional_params.py @@ -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