From 9b5693a5a287e557180324132f1fff8c3f583215 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 9 Aug 2024 11:30:34 -0700 Subject: [PATCH 1/4] ui allow adding cohere models --- ui/litellm-dashboard/src/components/model_dashboard.tsx | 2 ++ 1 file changed, 2 insertions(+) diff --git a/ui/litellm-dashboard/src/components/model_dashboard.tsx b/ui/litellm-dashboard/src/components/model_dashboard.tsx index 2d4fbe36f5b..d0d4b66a09e 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard.tsx @@ -147,6 +147,7 @@ enum Providers { MistralAI = "Mistral AI", OpenAI_Compatible = "OpenAI-Compatible Endpoints (Together AI, etc.)", Vertex_AI = "Vertex AI (Anthropic, Gemini, etc.)", + Cohere = "Cohere", Databricks = "Databricks", Ollama = "Ollama", } @@ -160,6 +161,7 @@ const provider_map: Record = { Bedrock: "bedrock", Groq: "groq", MistralAI: "mistral", + Cohere: "cohere_chat", OpenAI_Compatible: "openai", Vertex_AI: "vertex_ai", Databricks: "databricks", From df024fbbbc203b663ff46d81a496496815bff740 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 9 Aug 2024 12:08:25 -0700 Subject: [PATCH 2/4] add testing for cohere embeddings --- litellm/tests/test_embedding.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/litellm/tests/test_embedding.py b/litellm/tests/test_embedding.py index 9c1e4aa2cd6..31268395f14 100644 --- a/litellm/tests/test_embedding.py +++ b/litellm/tests/test_embedding.py @@ -282,19 +282,19 @@ async def test_cohere_embedding(sync_mode): # test_cohere_embedding() -def test_cohere_embedding3(): +@pytest.mark.parametrize("custom_llm_provider", ["cohere", "cohere_chat"]) +@pytest.mark.asyncio() +async def test_cohere_embedding3(custom_llm_provider): try: litellm.set_verbose = True - response = embedding( - model="embed-english-v3.0", + response = await litellm.aembedding( + model=f"{custom_llm_provider}/embed-english-v3.0", input=["good morning from litellm", "this is another item"], + timeout=None, + max_retries=0, ) print(f"response:", response) - custom_llm_provider = response._hidden_params["custom_llm_provider"] - - assert custom_llm_provider == "cohere" - except Exception as e: pytest.fail(f"Error occurred: {e}") From e734568b5a254fd28c4f8626e20e75086b3e6b85 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 9 Aug 2024 12:10:02 -0700 Subject: [PATCH 3/4] fix cohere / cohere_chat when timeout is None --- litellm/main.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index 1840a900a60..6200d8cb18a 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -3387,7 +3387,7 @@ def embedding( client=client, aembedding=aembedding, ) - elif custom_llm_provider == "cohere": + elif custom_llm_provider == "cohere" or custom_llm_provider == "cohere_chat": cohere_key = ( api_key or litellm.cohere_key @@ -3404,7 +3404,7 @@ def embedding( logging_obj=logging, model_response=EmbeddingResponse(), aembedding=aembedding, - timeout=float(timeout), + timeout=timeout, client=client, ) elif custom_llm_provider == "huggingface": From 05a4555fe73bb727d5443b3f648e585f8de1d6f1 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Fri, 9 Aug 2024 12:11:02 -0700 Subject: [PATCH 4/4] ui add cohere embedding models --- .../src/components/model_dashboard.tsx | 21 ++++++++++++++++++- 1 file changed, 20 insertions(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/model_dashboard.tsx b/ui/litellm-dashboard/src/components/model_dashboard.tsx index d0d4b66a09e..de92805b0bf 100644 --- a/ui/litellm-dashboard/src/components/model_dashboard.tsx +++ b/ui/litellm-dashboard/src/components/model_dashboard.tsx @@ -930,7 +930,26 @@ const ModelDashboard: React.FC = ({ _providerModels.push(key); } }); + + // Special case for cohere_chat + // we need both cohere_chat and cohere models to show on dropdown + if (providerKey == Providers.Cohere) { + console.log("adding cohere chat model") + Object.entries(modelMap).forEach(([key, value]) => { + if ( + value !== null && + typeof value === "object" && + "litellm_provider" in (value as object) && + ((value as any)["litellm_provider"] === "cohere") + ) { + _providerModels.push(key); + } + }); + } } + + + setProviderModels(_providerModels); console.log(`providerModels: ${providerModels}`); } @@ -1787,7 +1806,7 @@ const ModelDashboard: React.FC = ({