mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge pull request #38670 from BerriAI/devin_ai_38659_cohere_embed_dispatch
fix(bedrock): route all cohere.embed models to the cohere embedding config
This commit is contained in:
commit
9ed7de6c02
4 changed files with 52 additions and 3 deletions
|
|
@ -20,7 +20,9 @@ class BedrockCohereEmbeddingConfig:
|
|||
def map_openai_params(self, non_default_params: dict, optional_params: dict) -> dict:
|
||||
for k, v in non_default_params.items():
|
||||
if k == "encoding_format":
|
||||
optional_params["embedding_types"] = v if isinstance(v, list) else [v]
|
||||
optional_params["embedding_types"] = [
|
||||
"float" if fmt == "base64" else fmt for fmt in (tuple(v) if isinstance(v, list) else (v,))
|
||||
]
|
||||
elif k == "dimensions":
|
||||
optional_params["output_dimension"] = v
|
||||
return optional_params
|
||||
|
|
|
|||
|
|
@ -3581,7 +3581,7 @@ def get_optional_params_embeddings(
|
|||
object = litellm.AmazonTitanMultimodalEmbeddingG1Config()
|
||||
elif "amazon.titan-embed-text-v2:0" in model:
|
||||
object = litellm.AmazonTitanV2Config()
|
||||
elif "cohere.embed-multilingual-v3" in model or "cohere.embed-v4" in model:
|
||||
elif "cohere.embed" in model:
|
||||
object = litellm.BedrockCohereEmbeddingConfig()
|
||||
elif "twelvelabs" in model or "marengo" in model:
|
||||
object = litellm.TwelveLabsMarengoEmbeddingConfig()
|
||||
|
|
|
|||
|
|
@ -945,7 +945,7 @@ def test_titan_image_embedding_cost_uses_per_image_rate():
|
|||
"encoding_format,expected_embedding_types",
|
||||
[
|
||||
("float", ["float"]),
|
||||
("base64", ["base64"]),
|
||||
("base64", ["float"]),
|
||||
(["float", "int8"], ["float", "int8"]),
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4432,6 +4432,53 @@ class TestVertexEmbeddingEncodingFormat:
|
|||
assert optional_params.get("outputDimensionality") == 256
|
||||
|
||||
|
||||
class TestBedrockCohereEmbeddingDispatch:
|
||||
"""All bedrock cohere.embed models must route to BedrockCohereEmbeddingConfig,
|
||||
not just multilingual-v3/v4: english-v3 was falling into the unmapped
|
||||
else-branch and rejecting encoding_format. Issue #38659."""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"cohere.embed-english-v3",
|
||||
"cohere.embed-multilingual-v3",
|
||||
"cohere.embed-v4:0",
|
||||
],
|
||||
)
|
||||
def test_cohere_embed_models_accept_encoding_format(self, model):
|
||||
optional_params = litellm.utils.get_optional_params_embeddings(
|
||||
model=model,
|
||||
encoding_format="float",
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
assert optional_params.get("embedding_types") == ["float"]
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
"cohere.embed-english-v3",
|
||||
"cohere.embed-multilingual-v3",
|
||||
"cohere.embed-v4:0",
|
||||
],
|
||||
)
|
||||
def test_cohere_embed_models_map_base64_to_float(self, model):
|
||||
optional_params = litellm.utils.get_optional_params_embeddings(
|
||||
model=model,
|
||||
encoding_format="base64",
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
assert optional_params.get("embedding_types") == ["float"]
|
||||
|
||||
def test_cohere_embed_english_v3_maps_dimensions(self):
|
||||
optional_params = litellm.utils.get_optional_params_embeddings(
|
||||
model="cohere.embed-english-v3",
|
||||
encoding_format="float",
|
||||
dimensions=512,
|
||||
custom_llm_provider="bedrock",
|
||||
)
|
||||
assert optional_params.get("output_dimension") == 512
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model",
|
||||
[
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue