diff --git a/litellm/main.py b/litellm/main.py index 46a024a7631..6cdba2e8e25 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -85,6 +85,10 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( from litellm.llms.base_llm import BaseConfig, BaseImageGenerationConfig from litellm.llms.bedrock.common_utils import BedrockModelInfo from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.llms.vertex_ai.common_utils import ( + VertexAIModelRoute, + get_vertex_ai_model_route, +) from litellm.realtime_api.main import _realtime_health_check from litellm.secret_managers.main import get_secret_bool, get_secret_str from litellm.types.router import GenericLiteLLMParams @@ -150,7 +154,6 @@ from .llms.bedrock.chat import BedrockConverseLLM, BedrockLLM from .llms.bedrock.embed.embedding import BedrockEmbedding from .llms.bedrock.image.image_handler import BedrockImageGeneration from .llms.bytez.chat.transformation import BytezChatConfig -from .llms.lemonade.chat.transformation import LemonadeChatConfig from .llms.codestral.completion.handler import CodestralTextCompletion from .llms.cohere.embed import handler as cohere_embed from .llms.custom_httpx.aiohttp_handler import BaseLLMAIOHTTPHandler @@ -162,6 +165,7 @@ from .llms.gemini.common_utils import get_api_key_from_env from .llms.groq.chat.handler import GroqChatCompletion from .llms.heroku.chat.transformation import HerokuChatConfig from .llms.huggingface.embedding.handler import HuggingFaceEmbedding +from .llms.lemonade.chat.transformation import LemonadeChatConfig from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion from .llms.oci.chat.transformation import OCIChatConfig from .llms.ollama.completion import handler as ollama @@ -192,6 +196,7 @@ from .llms.vertex_ai.multimodal_embeddings.embedding_handler import ( from .llms.vertex_ai.text_to_speech.text_to_speech_handler import VertexTextToSpeechAPI from .llms.vertex_ai.vertex_ai_partner_models.main import VertexAIPartnerModels from .llms.vertex_ai.vertex_embeddings.embedding_handler import VertexEmbedding +from .llms.vertex_ai.vertex_gemma_models.main import VertexAIGemmaModels from .llms.vertex_ai.vertex_model_garden.main import VertexAIModelGardenModels from .llms.vllm.completion import handler as vllm_handler from .llms.watsonx.chat.handler import WatsonXChatHandler @@ -255,6 +260,7 @@ vertex_multimodal_embedding = VertexMultimodalEmbedding() vertex_image_generation = VertexImageGeneration() google_batch_embeddings = GoogleBatchEmbeddings() vertex_partner_models_chat_completion = VertexAIPartnerModels() +vertex_gemma_chat_completion = VertexAIGemmaModels() vertex_model_garden_chat_completion = VertexAIModelGardenModels() vertex_text_to_speech = VertexTextToSpeechAPI() sagemaker_llm = SagemakerLLM() @@ -2875,7 +2881,7 @@ def completion( # type: ignore # noqa: PLR0915 extra_headers=headers, ) - elif custom_llm_provider == "vertex_ai": + elif custom_llm_provider == "vertex_ai": vertex_ai_project = ( optional_params.pop("vertex_project", None) or optional_params.pop("vertex_ai_project", None) @@ -2897,7 +2903,9 @@ def completion( # type: ignore # noqa: PLR0915 api_base = api_base or litellm.api_base or get_secret("VERTEXAI_API_BASE") new_params = safe_deep_copy(optional_params or {}) - if vertex_partner_models_chat_completion.is_vertex_partner_model(model): + model_route = get_vertex_ai_model_route(model=model, litellm_params=litellm_params) + + if model_route == VertexAIModelRoute.PARTNER_MODELS: model_response = vertex_partner_models_chat_completion.completion( model=model, messages=messages, @@ -2918,10 +2926,7 @@ def completion( # type: ignore # noqa: PLR0915 timeout=timeout, client=client, ) - elif "gemini" in model or ( - litellm_params.get("base_model") is not None - and "gemini" in litellm_params["base_model"] - ): + elif model_route == VertexAIModelRoute.GEMINI: model_response = vertex_chat_completion.completion( # type: ignore model=model, messages=messages, @@ -2943,7 +2948,29 @@ def completion( # type: ignore # noqa: PLR0915 api_base=api_base, extra_headers=headers, ) - elif "openai" in model: + elif model_route == VertexAIModelRoute.GEMMA: + # Vertex Gemma Models with custom prediction endpoint + model_response = vertex_gemma_chat_completion.completion( + model=model, + messages=messages, + model_response=model_response, + print_verbose=print_verbose, + optional_params=new_params, + litellm_params=litellm_params, # type: ignore + logger_fn=logger_fn, + encoding=encoding, + api_base=api_base, + vertex_location=vertex_ai_location, + vertex_project=vertex_ai_project, + vertex_credentials=vertex_credentials, + logging_obj=logging, + acompletion=acompletion, + headers=headers, + custom_prompt_dict=custom_prompt_dict, + timeout=timeout, + client=client, + ) + elif model_route == VertexAIModelRoute.MODEL_GARDEN: # Vertex Model Garden - OpenAI compatible models model_response = vertex_model_garden_chat_completion.completion( model=model, @@ -2965,7 +2992,7 @@ def completion( # type: ignore # noqa: PLR0915 timeout=timeout, client=client, ) - else: + else: # VertexAIModelRoute.NON_GEMINI model_response = vertex_ai_non_gemini.completion( model=model, messages=messages,