Merge pull request #5136 from BerriAI/litellm_add_cohere_embedding_models

ui allow adding cohere models
This commit is contained in:
Ishaan Jaff 2024-08-09 12:19:19 -07:00 • committed by GitHub
commit db7897dcc0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 31 additions and 10 deletions

View file

@ -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":

View file

@ -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}")

View file

@ -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<string, string> = {
Bedrock: "bedrock",
Groq: "groq",
MistralAI: "mistral",
Cohere: "cohere_chat",
OpenAI_Compatible: "openai",
Vertex_AI: "vertex_ai",
Databricks: "databricks",
@ -928,7 +930,26 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
_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}`);
}
@ -1785,7 +1806,7 @@ const ModelDashboard: React.FC<ModelDashboardProps> = ({
</Row>
<Form.Item
label="LiteLLM Model Name(s)"
tooltip="Actual model name used for making litellm.completion() call."
tooltip="Actual model name used for making litellm.completion() / litellm.embedding() call."
className="mb-0"
>
<Form.Item