From 3469b5b911bd0a0fed3179975abd630753cf11c5 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 8 Jan 2024 07:38:55 +0530 Subject: [PATCH] fix(utils.py): map optional params for gemini --- litellm/tests/test_google_ai_studio_gemini.py | 2 +- litellm/utils.py | 7 +++++-- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/litellm/tests/test_google_ai_studio_gemini.py b/litellm/tests/test_google_ai_studio_gemini.py index ff24933cc9..7cebd25372 100644 --- a/litellm/tests/test_google_ai_studio_gemini.py +++ b/litellm/tests/test_google_ai_studio_gemini.py @@ -25,7 +25,7 @@ def generate_text(): ] } ] - response = litellm.completion(model="gemini/gemini-pro-vision", messages=messages) + response = litellm.completion(model="gemini/gemini-pro-vision", messages=messages, stop="Hello world") print(response) assert isinstance(response.choices[0].message.content, str) == True except Exception as exception: diff --git a/litellm/utils.py b/litellm/utils.py index b9f252efb1..ff09390dad 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -3435,7 +3435,7 @@ def get_optional_params( if presence_penalty is not None: optional_params["presencePenalty"] = {"scale": presence_penalty} elif ( - custom_llm_provider == "palm" + custom_llm_provider == "palm" or custom_llm_provider == "gemini" ): # https://developers.generativeai.google/tutorials/curl_quickstart ## check if unsupported param passed in supported_params = ["temperature", "top_p", "stream", "n", "stop", "max_tokens"] @@ -3450,7 +3450,10 @@ def get_optional_params( if n is not None: optional_params["candidate_count"] = n if stop is not None: - optional_params["stop_sequences"] = stop + if isinstance(stop, str): + optional_params["stop_sequences"] = [stop] + elif isinstance(stop, list): + optional_params["stop_sequences"] = stop if max_tokens is not None: optional_params["max_output_tokens"] = max_tokens elif custom_llm_provider == "vertex_ai":