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:
kartik-mem0
2026-07-09 21:30:39 +05:30
parent 99206f0c64
commit 493a62cb03
2 changed files with 41 additions and 6 deletions
+11 -2
View File
@@ -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)
+30 -4
View File
@@ -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