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) <noreply@anthropic.com>
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
+18
-10
@@ -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
|
||||
|
||||
+2
-23
@@ -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]:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user