fix(embeddings): size Vertex AI batches by model family
gemini-embedding-* accepts one input text per predict request, but embed_batch() chunked at 250 for every model while defaulting to gemini-embedding-001. Derive the chunk size from the model instead.
This commit is contained in:
@@ -8,6 +8,14 @@ from mem0.embeddings.base import EmbeddingBase
|
||||
from mem0.utils.gcp_auth import GCPAuthenticator
|
||||
|
||||
|
||||
def _max_instances_per_request(model: str) -> int:
|
||||
# Vertex AI's per-request input cap depends on the model family: gemini-embedding-*
|
||||
# accepts exactly one input text per predict() call, while text-embedding-004/005 and
|
||||
# text-multilingual-embedding-002 accept up to 250.
|
||||
# https://docs.cloud.google.com/vertex-ai/docs/samples/aiplatform-sdk-embedding
|
||||
return 1 if model.startswith("gemini-embedding") else 250
|
||||
|
||||
|
||||
class VertexAIEmbedding(EmbeddingBase):
|
||||
def __init__(self, config: Optional[BaseEmbedderConfig] = None):
|
||||
super().__init__(config)
|
||||
@@ -72,8 +80,9 @@ class VertexAIEmbedding(EmbeddingBase):
|
||||
raise ValueError(f"Invalid memory action: {memory_action}")
|
||||
embedding_type = self.embedding_types[memory_action]
|
||||
all_embeddings = []
|
||||
for i in range(0, len(texts), 250):
|
||||
chunk = texts[i : i + 250]
|
||||
chunk_size = _max_instances_per_request(self.config.model)
|
||||
for i in range(0, len(texts), chunk_size):
|
||||
chunk = texts[i : i + chunk_size]
|
||||
inputs = [TextEmbeddingInput(text=t, task_type=embedding_type) for t in chunk]
|
||||
results = self.model.get_embeddings(texts=inputs, output_dimensionality=self.config.embedding_dims)
|
||||
all_embeddings.extend(r.values for r in results)
|
||||
|
||||
@@ -163,7 +163,9 @@ def test_invalid_memory_action(mock_text_embedding_model, mock_config):
|
||||
|
||||
@patch("mem0.embeddings.vertexai.TextEmbeddingModel")
|
||||
def test_embed_batch_single_call(mock_text_embedding_model, mock_os_environ, mock_config):
|
||||
mock_config.return_value.model = "gemini-embedding-001"
|
||||
# text-embedding-005 accepts up to 250 inputs per request, so 2 texts fit in one call.
|
||||
# (gemini-embedding-* only accepts 1 input per request; see test_embed_batch_default_model_sends_one_text_per_call.)
|
||||
mock_config.return_value.model = "text-embedding-005"
|
||||
mock_config.return_value.embedding_dims = 256
|
||||
|
||||
config = mock_config()
|
||||
@@ -198,7 +200,9 @@ def test_embed_batch_empty_list(mock_text_embedding_model, mock_os_environ, mock
|
||||
|
||||
@patch("mem0.embeddings.vertexai.TextEmbeddingModel")
|
||||
def test_embed_batch_count_mismatch_raises(mock_text_embedding_model, mock_os_environ, mock_config):
|
||||
mock_config.return_value.model = "gemini-embedding-001"
|
||||
# Use a 250-cap model so both texts land in a single chunk/call, keeping the mismatch
|
||||
# (1 embedding returned for that one call) triggered by that single call.
|
||||
mock_config.return_value.model = "text-embedding-005"
|
||||
mock_config.return_value.embedding_dims = 256
|
||||
|
||||
config = mock_config()
|
||||
@@ -261,8 +265,8 @@ def test_embed_batch_invalid_memory_action_raises(mock_text_embedding_model, moc
|
||||
|
||||
@patch("mem0.embeddings.vertexai.TextEmbeddingModel")
|
||||
def test_embed_batch_chunking_triggers_two_api_calls(mock_text_embedding_model, mock_os_environ, mock_config):
|
||||
"""300 texts must produce exactly 2 get_embeddings calls (chunks of 250 and 50)."""
|
||||
mock_config.return_value.model = "gemini-embedding-001"
|
||||
"""text-embedding-005 accepts up to 250 inputs/request: 300 texts -> 2 calls (chunks of 250 and 50)."""
|
||||
mock_config.return_value.model = "text-embedding-005"
|
||||
mock_config.return_value.embedding_dims = 256
|
||||
|
||||
config = mock_config()
|
||||
@@ -278,3 +282,25 @@ def test_embed_batch_chunking_triggers_two_api_calls(mock_text_embedding_model,
|
||||
|
||||
assert mock_text_embedding_model.from_pretrained.return_value.get_embeddings.call_count == 2
|
||||
assert len(result) == 300
|
||||
|
||||
|
||||
@patch("mem0.embeddings.vertexai.TextEmbeddingModel")
|
||||
def test_embed_batch_default_model_sends_one_text_per_call(mock_text_embedding_model, mock_os_environ, mock_config):
|
||||
"""gemini-embedding-001 (default) accepts 1 input/request: 300 texts -> 300 calls of 1 text each."""
|
||||
mock_config.return_value.model = "gemini-embedding-001"
|
||||
mock_config.return_value.embedding_dims = 256
|
||||
|
||||
config = mock_config()
|
||||
embedder = VertexAIEmbedding(config)
|
||||
|
||||
def make_single_response(texts, output_dimensionality):
|
||||
assert len(texts) == 1
|
||||
return [Mock(values=[0.1, 0.2]) for _ in texts]
|
||||
|
||||
mock_text_embedding_model.from_pretrained.return_value.get_embeddings.side_effect = make_single_response
|
||||
|
||||
texts = [f"text {i}" for i in range(300)]
|
||||
result = embedder.embed_batch(texts)
|
||||
|
||||
assert mock_text_embedding_model.from_pretrained.return_value.get_embeddings.call_count == 300
|
||||
assert len(result) == 300
|
||||
|
||||
Reference in New Issue
Block a user