From 3402c678059e933b38446b36912db92542bb94bb Mon Sep 17 00:00:00 2001 From: Soumil Rathi Date: Mon, 13 Apr 2026 09:01:07 -0700 Subject: [PATCH] fix: address PR review feedback - Remove deprecated custom_update_memory_prompt from MemoryConfig entirely - Add MAX_BATCH=100 chunking guard to embed_batch in openai.py and azure_openai.py - Move all inline imports (mem0.utils.*) to top-level in main.py - Remove custom_update_memory_prompt from test fixture Co-Authored-By: Claude Opus 4.6 (1M context) --- mem0/configs/base.py | 4 ---- mem0/embeddings/azure_openai.py | 20 ++++++++++++++------ mem0/embeddings/openai.py | 28 ++++++++++++++++++---------- mem0/memory/main.py | 25 ++----------------------- tests/test_main.py | 1 - 5 files changed, 34 insertions(+), 44 deletions(-) diff --git a/mem0/configs/base.py b/mem0/configs/base.py index 29568563e..52a0178da 100644 --- a/mem0/configs/base.py +++ b/mem0/configs/base.py @@ -60,10 +60,6 @@ class MemoryConfig(BaseModel): description="Custom instructions for fact extraction", default=None, ) - custom_update_memory_prompt: Optional[str] = Field( - description="Custom prompt for the update memory (deprecated: use custom_instructions)", - default=None, - ) class AzureConfig(BaseModel): diff --git a/mem0/embeddings/azure_openai.py b/mem0/embeddings/azure_openai.py index 7818de481..95670a8a1 100644 --- a/mem0/embeddings/azure_openai.py +++ b/mem0/embeddings/azure_openai.py @@ -55,10 +55,18 @@ class AzureOpenAIEmbedding(EmbeddingBase): return self.client.embeddings.create(input=[text], model=self.config.model).data[0].embedding def embed_batch(self, texts, memory_action="add"): - """Embed multiple texts in a single Azure OpenAI API call.""" + """Embed multiple texts in a single Azure OpenAI API call. + + Automatically chunks into batches of 100 to stay within API limits. + """ + MAX_BATCH = 100 texts = [text.replace("\n", " ") for text in texts] - response = self.client.embeddings.create( - input=texts, - model=self.config.model, - ) - return [item.embedding for item in sorted(response.data, key=lambda x: x.index)] + all_embeddings = [] + for i in range(0, len(texts), MAX_BATCH): + chunk = texts[i : i + MAX_BATCH] + response = self.client.embeddings.create( + input=chunk, + model=self.config.model, + ) + all_embeddings.extend(item.embedding for item in sorted(response.data, key=lambda x: x.index)) + return all_embeddings diff --git a/mem0/embeddings/openai.py b/mem0/embeddings/openai.py index bbf2b3456..cede4f1cb 100644 --- a/mem0/embeddings/openai.py +++ b/mem0/embeddings/openai.py @@ -55,14 +55,22 @@ class OpenAIEmbedding(EmbeddingBase): return self.client.embeddings.create(**kwargs).data[0].embedding def embed_batch(self, texts, memory_action="add"): - """Embed multiple texts in a single OpenAI API call.""" + """Embed multiple texts in a single OpenAI API call. + + Automatically chunks into batches of 100 to stay within API limits. + """ + MAX_BATCH = 100 texts = [text.replace("\n", " ") for text in texts] - kwargs = { - "input": texts, - "model": self.config.model, - "encoding_format": "float", - } - if self._pass_dimensions_to_api: - kwargs["dimensions"] = self.config.embedding_dims - response = self.client.embeddings.create(**kwargs) - return [item.embedding for item in sorted(response.data, key=lambda x: x.index)] + all_embeddings = [] + for i in range(0, len(texts), MAX_BATCH): + chunk = texts[i : i + MAX_BATCH] + kwargs = { + "input": chunk, + "model": self.config.model, + "encoding_format": "float", + } + if self._pass_dimensions_to_api: + kwargs["dimensions"] = self.config.embedding_dims + response = self.client.embeddings.create(**kwargs) + all_embeddings.extend(item.embedding for item in sorted(response.data, key=lambda x: x.index)) + return all_embeddings diff --git a/mem0/memory/main.py b/mem0/memory/main.py index ba6226ffd..f7e7a5af6 100644 --- a/mem0/memory/main.py +++ b/mem0/memory/main.py @@ -24,7 +24,8 @@ from mem0.configs.prompts import ( get_update_memory_messages, ) from mem0.utils.lemmatization import lemmatize_for_bm25 -from mem0.utils.entity_extraction import extract_entities_batch +from mem0.utils.entity_extraction import extract_entities, extract_entities_batch +from mem0.utils.scoring import ENTITY_BOOST_WEIGHT, get_bm25_params, normalize_bm25, score_and_rank from mem0.exceptions import ValidationError as Mem0ValidationError from mem0.memory.base import MemoryBase from mem0.memory.setup import mem0_dir, setup_config @@ -191,7 +192,6 @@ class Memory(MemoryBase): def __init__(self, config: MemoryConfig = MemoryConfig()): self.config = config - self.custom_update_memory_prompt = self.config.custom_update_memory_prompt self.embedding_model = EmbedderFactory.create( self.config.embedder.provider, self.config.embedder.config, @@ -217,15 +217,6 @@ class Memory(MemoryBase): # Entity store is initialized lazily on first use self._entity_store = None - if self.config.custom_update_memory_prompt: - import warnings - warnings.warn( - "custom_update_memory_prompt is deprecated and has no effect in the v3 pipeline. " - "Use custom_instructions instead.", - DeprecationWarning, - stacklevel=2, - ) - if self.config.graph_store.config: provider = self.config.graph_store.provider self.graph = GraphStoreFactory.create(provider, self.config) @@ -1145,10 +1136,6 @@ class Memory(MemoryBase): return False def _search_vector_store(self, query, filters, limit, threshold=0.1): - from mem0.utils.lemmatization import lemmatize_for_bm25 - from mem0.utils.entity_extraction import extract_entities - from mem0.utils.scoring import get_bm25_params, normalize_bm25, score_and_rank, ENTITY_BOOST_WEIGHT - # Guard against None threshold (backward compat) if threshold is None: threshold = 0.1 @@ -1261,8 +1248,6 @@ class Memory(MemoryBase): Returns: Dict mapping memory_id (str) -> max entity boost [0, 0.5]. """ - from mem0.utils.scoring import ENTITY_BOOST_WEIGHT - # Deduplicate entities (max 8) seen = set() deduped = [] @@ -2494,10 +2479,6 @@ class AsyncMemory(MemoryBase): return False async def _search_vector_store(self, query, filters, limit, threshold=0.1): - from mem0.utils.lemmatization import lemmatize_for_bm25 - from mem0.utils.entity_extraction import extract_entities - from mem0.utils.scoring import get_bm25_params, normalize_bm25, score_and_rank, ENTITY_BOOST_WEIGHT - if threshold is None: threshold = 0.1 @@ -2596,8 +2577,6 @@ class AsyncMemory(MemoryBase): async def _compute_entity_boosts_async(self, query_entities, filters): """Async version of entity boost computation.""" - from mem0.utils.scoring import ENTITY_BOOST_WEIGHT - seen = set() deduped = [] for entity_type, entity_text in query_entities[:8]: diff --git a/tests/test_main.py b/tests/test_main.py index b47391f12..a8a61bdff 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -61,7 +61,6 @@ def memory_custom_instance(): config = MemoryConfig( version="v1.1", custom_instructions="custom prompt extracting memory in json format", - custom_update_memory_prompt="custom prompt determining memory update", ) config.graph_store.config = {"some_config": "value"} return Memory(config)