From 493a62cb03ec522c06c8e63ae2764ebc65faeaf2 Mon Sep 17 00:00:00 2001 From: kartik-mem0 Date: Thu, 9 Jul 2026 21:30:39 +0530 Subject: [PATCH] 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. --- mem0/embeddings/vertexai.py | 13 ++++++-- tests/embeddings/test_vertexai_embeddings.py | 34 +++++++++++++++++--- 2 files changed, 41 insertions(+), 6 deletions(-) diff --git a/mem0/embeddings/vertexai.py b/mem0/embeddings/vertexai.py index 039bfd5f8..ae75c14f5 100644 --- a/mem0/embeddings/vertexai.py +++ b/mem0/embeddings/vertexai.py @@ -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) diff --git a/tests/embeddings/test_vertexai_embeddings.py b/tests/embeddings/test_vertexai_embeddings.py index b78bff66e..ebb9c1d7f 100644 --- a/tests/embeddings/test_vertexai_embeddings.py +++ b/tests/embeddings/test_vertexai_embeddings.py @@ -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