From 341f88d1917db416b31f8fd09db8a3a85c561e30 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 11 Jul 2024 12:59:42 -0700 Subject: [PATCH 1/2] fix supports vision --- litellm/proxy/proxy_config.yaml | 3 +++ litellm/types/utils.py | 1 + litellm/utils.py | 11 ++++++----- 3 files changed, 10 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 5fd4e5b62c2..762155ed987 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -4,6 +4,9 @@ model_list: model: openai/fake api_key: fake-key api_base: https://exampleopenaiendpoint-production.up.railway.app/ + - model_name: gemini-flash + litellm_params: + model: gemini/gemini-1.5-flash general_settings: master_key: sk-1234 diff --git a/litellm/types/utils.py b/litellm/types/utils.py index a90f93484f1..9c32fbf4d61 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -73,6 +73,7 @@ class ModelInfo(TypedDict, total=False): supported_openai_params: Required[Optional[List[str]]] supports_system_messages: Optional[bool] supports_response_schema: Optional[bool] + supports_vision: Optional[bool] class GenericStreamingChunk(TypedDict): diff --git a/litellm/utils.py b/litellm/utils.py index cf2c679a844..2dab185a3d5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4829,6 +4829,7 @@ def get_model_info(model: str, custom_llm_provider: Optional[str] = None) -> Mod supports_response_schema=_model_info.get( "supports_response_schema", None ), + supports_vision=_model_info.get("supports_vision", None), ) except Exception: raise Exception( @@ -8126,7 +8127,7 @@ class CustomStreamWrapper: if chunk.startswith(self.complete_response): # Remove last_sent_chunk only if it appears at the start of the new chunk - chunk = chunk[len(self.complete_response):] + chunk = chunk[len(self.complete_response) :] self.complete_response += chunk return chunk @@ -10124,7 +10125,7 @@ def mock_completion_streaming_obj( model_response, mock_response, model, n: Optional[int] = None ): for i in range(0, len(mock_response), 3): - completion_obj = Delta(role="assistant", content=mock_response[i: i + 3]) + completion_obj = Delta(role="assistant", content=mock_response[i : i + 3]) if n is None: model_response.choices[0].delta = completion_obj else: @@ -10133,7 +10134,7 @@ def mock_completion_streaming_obj( _streaming_choice = litellm.utils.StreamingChoices( index=j, delta=litellm.utils.Delta( - role="assistant", content=mock_response[i: i + 3] + role="assistant", content=mock_response[i : i + 3] ), ) _all_choices.append(_streaming_choice) @@ -10145,7 +10146,7 @@ async def async_mock_completion_streaming_obj( model_response, mock_response, model, n: Optional[int] = None ): for i in range(0, len(mock_response), 3): - completion_obj = Delta(role="assistant", content=mock_response[i: i + 3]) + completion_obj = Delta(role="assistant", content=mock_response[i : i + 3]) if n is None: model_response.choices[0].delta = completion_obj else: @@ -10154,7 +10155,7 @@ async def async_mock_completion_streaming_obj( _streaming_choice = litellm.utils.StreamingChoices( index=j, delta=litellm.utils.Delta( - role="assistant", content=mock_response[i: i + 3] + role="assistant", content=mock_response[i : i + 3] ), ) _all_choices.append(_streaming_choice) From 46493303edc301a514daaa9755973e90dd36bc2d Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 11 Jul 2024 13:04:18 -0700 Subject: [PATCH 2/2] test get mode info for gemini/gemini-1.5-flash --- litellm/tests/test_get_model_info.py | 13 +++++++++++-- 1 file changed, 11 insertions(+), 2 deletions(-) diff --git a/litellm/tests/test_get_model_info.py b/litellm/tests/test_get_model_info.py index 3fd6a6d22f6..687aa062f15 100644 --- a/litellm/tests/test_get_model_info.py +++ b/litellm/tests/test_get_model_info.py @@ -1,13 +1,16 @@ # What is this? ## Unit testing for the 'get_model_info()' function -import os, sys, traceback +import os +import sys +import traceback sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path +import pytest + import litellm from litellm import get_model_info -import pytest def test_get_model_info_simple_model_name(): @@ -37,3 +40,9 @@ def test_get_model_info_custom_llm_with_same_name_vllm(): pytest.fail("Expected get model info to fail for an unmapped model/provider") except Exception: pass + + +def test_get_model_info_shows_correct_supports_vision(): + info = litellm.get_model_info("gemini/gemini-1.5-flash") + print("info", info) + assert info["supports_vision"] is True