Fix: prevent double embedding in mem0.add (fixes #3723)
This fix addresses issue #3723 where mem0.add() was calling the embedding API twice, unnecessarily doubling costs and latency. Changes made: 1. Modified _create_memory() to accept embeddings as either a dict (for caching) or a precomputed vector, preventing redundant calls 2. Updated infer=False path to pass embeddings as a dict 3. Added caching for action_text embeddings in infer=True path for both ADD and UPDATE operations, since the LLM may rephrase facts 4. Applied same fixes to both sync and async Memory classes 5. Added regression test to verify embedding is called only once The root cause was that when infer=False, embeddings were passed directly to _create_memory without a dict wrapper, causing it to re-embed. When infer=True, if the LLM rephrased extracted facts, the action_text wouldn't match the cache key, triggering re-embedding.
This commit is contained in:
committed by
kartik-mem0
parent
7cebaba0a2
commit
211a7570e7
+34
-6
@@ -479,7 +479,8 @@ class Memory(MemoryBase):
|
||||
|
||||
msg_content = message_dict["content"]
|
||||
msg_embeddings = self.embedding_model.embed(msg_content, "add")
|
||||
mem_id = self._create_memory(msg_content, msg_embeddings, per_msg_meta)
|
||||
# Pass embeddings as a dict so _create_memory can reuse the cached embedding
|
||||
mem_id = self._create_memory(msg_content, {msg_content: msg_embeddings}, per_msg_meta)
|
||||
|
||||
returned_memories.append(
|
||||
{
|
||||
@@ -608,6 +609,9 @@ class Memory(MemoryBase):
|
||||
|
||||
event_type = resp.get("event")
|
||||
if event_type == "ADD":
|
||||
# Ensure action_text has an embedding cached to avoid redundant API calls
|
||||
if action_text not in new_message_embeddings:
|
||||
new_message_embeddings[action_text] = self.embedding_model.embed(action_text, "add")
|
||||
memory_id = self._create_memory(
|
||||
data=action_text,
|
||||
existing_embeddings=new_message_embeddings,
|
||||
@@ -615,6 +619,9 @@ class Memory(MemoryBase):
|
||||
)
|
||||
returned_memories.append({"id": memory_id, "memory": action_text, "event": event_type})
|
||||
elif event_type == "UPDATE":
|
||||
# Ensure action_text has an embedding cached to avoid redundant API calls
|
||||
if action_text not in new_message_embeddings:
|
||||
new_message_embeddings[action_text] = self.embedding_model.embed(action_text, "update")
|
||||
self._update_memory(
|
||||
memory_id=temp_uuid_mapping[resp.get("id")],
|
||||
data=action_text,
|
||||
@@ -1155,8 +1162,11 @@ class Memory(MemoryBase):
|
||||
|
||||
def _create_memory(self, data, existing_embeddings, metadata=None):
|
||||
logger.debug(f"Creating memory with {data=}")
|
||||
if data in existing_embeddings:
|
||||
# existing_embeddings may be a dict (preferred) or a precomputed vector
|
||||
if isinstance(existing_embeddings, dict) and data in existing_embeddings:
|
||||
embeddings = existing_embeddings[data]
|
||||
elif isinstance(existing_embeddings, list):
|
||||
embeddings = existing_embeddings
|
||||
else:
|
||||
embeddings = self.embedding_model.embed(data, memory_action="add")
|
||||
memory_id = str(uuid.uuid4())
|
||||
@@ -1253,8 +1263,10 @@ class Memory(MemoryBase):
|
||||
if "role" not in new_metadata and "role" in existing_memory.payload:
|
||||
new_metadata["role"] = existing_memory.payload["role"]
|
||||
|
||||
if data in existing_embeddings:
|
||||
if isinstance(existing_embeddings, dict) and data in existing_embeddings:
|
||||
embeddings = existing_embeddings[data]
|
||||
elif isinstance(existing_embeddings, list):
|
||||
embeddings = existing_embeddings
|
||||
else:
|
||||
embeddings = self.embedding_model.embed(data, "update")
|
||||
|
||||
@@ -1523,7 +1535,8 @@ class AsyncMemory(MemoryBase):
|
||||
|
||||
msg_content = message_dict["content"]
|
||||
msg_embeddings = await asyncio.to_thread(self.embedding_model.embed, msg_content, "add")
|
||||
mem_id = await self._create_memory(msg_content, msg_embeddings, per_msg_meta)
|
||||
# Pass embeddings as a dict so _create_memory can reuse the cached embedding
|
||||
mem_id = await self._create_memory(msg_content, {msg_content: msg_embeddings}, per_msg_meta)
|
||||
|
||||
returned_memories.append(
|
||||
{
|
||||
@@ -1651,6 +1664,11 @@ class AsyncMemory(MemoryBase):
|
||||
event_type = resp.get("event")
|
||||
|
||||
if event_type == "ADD":
|
||||
# Ensure action_text has an embedding cached to avoid redundant API calls
|
||||
if action_text not in new_message_embeddings:
|
||||
new_message_embeddings[action_text] = await asyncio.to_thread(
|
||||
self.embedding_model.embed, action_text, "add"
|
||||
)
|
||||
task = asyncio.create_task(
|
||||
self._create_memory(
|
||||
data=action_text,
|
||||
@@ -1660,6 +1678,11 @@ class AsyncMemory(MemoryBase):
|
||||
)
|
||||
memory_tasks.append((task, resp, "ADD", None))
|
||||
elif event_type == "UPDATE":
|
||||
# Ensure action_text has an embedding cached to avoid redundant API calls
|
||||
if action_text not in new_message_embeddings:
|
||||
new_message_embeddings[action_text] = await asyncio.to_thread(
|
||||
self.embedding_model.embed, action_text, "update"
|
||||
)
|
||||
task = asyncio.create_task(
|
||||
self._update_memory(
|
||||
memory_id=temp_uuid_mapping[resp["id"]],
|
||||
@@ -2231,8 +2254,11 @@ class AsyncMemory(MemoryBase):
|
||||
|
||||
async def _create_memory(self, data, existing_embeddings, metadata=None):
|
||||
logger.debug(f"Creating memory with {data=}")
|
||||
if data in existing_embeddings:
|
||||
# existing_embeddings may be a dict (preferred) or a precomputed vector
|
||||
if isinstance(existing_embeddings, dict) and data in existing_embeddings:
|
||||
embeddings = existing_embeddings[data]
|
||||
elif isinstance(existing_embeddings, list):
|
||||
embeddings = existing_embeddings
|
||||
else:
|
||||
embeddings = await asyncio.to_thread(self.embedding_model.embed, data, memory_action="add")
|
||||
|
||||
@@ -2349,8 +2375,10 @@ class AsyncMemory(MemoryBase):
|
||||
if "role" not in new_metadata and "role" in existing_memory.payload:
|
||||
new_metadata["role"] = existing_memory.payload["role"]
|
||||
|
||||
if data in existing_embeddings:
|
||||
if isinstance(existing_embeddings, dict) and data in existing_embeddings:
|
||||
embeddings = existing_embeddings[data]
|
||||
elif isinstance(existing_embeddings, list):
|
||||
embeddings = existing_embeddings
|
||||
else:
|
||||
embeddings = await asyncio.to_thread(self.embedding_model.embed, data, "update")
|
||||
|
||||
|
||||
@@ -438,3 +438,82 @@ async def test_async_update_nonexistent_memory_raises_error(mock_sqlite, mock_ll
|
||||
await memory._update_memory("non-existent-id", "new data", {"new data": [0.1, 0.2]})
|
||||
|
||||
mock_vector_store.update.assert_not_called()
|
||||
|
||||
|
||||
@patch('mem0.utils.factory.EmbedderFactory.create')
|
||||
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
||||
@patch('mem0.utils.factory.LlmFactory.create')
|
||||
@patch('mem0.memory.storage.SQLiteManager')
|
||||
def test_add_infer_false_embeds_once(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory):
|
||||
"""
|
||||
Regression test for issue #3723: adding with infer=False should not trigger duplicate embedding calls.
|
||||
|
||||
Root cause: _create_memory expected a dict for existing_embeddings but received a raw list[float],
|
||||
causing the cache check `data in existing_embeddings` to always fail and trigger a redundant embed.
|
||||
"""
|
||||
embedder = MagicMock()
|
||||
embedder.embed.return_value = [0.1, 0.2, 0.3]
|
||||
embedder.config = MagicMock(embedding_dims=3)
|
||||
mock_embedder_factory.return_value = embedder
|
||||
|
||||
mock_vector_store = MagicMock()
|
||||
mock_vector_store.search.return_value = []
|
||||
mock_vector_store.insert.return_value = None
|
||||
mock_vector_store.get.return_value = None
|
||||
telemetry_vector_store = MagicMock()
|
||||
mock_vector_factory.side_effect = [mock_vector_store, telemetry_vector_store]
|
||||
|
||||
mock_llm_factory.return_value = MagicMock()
|
||||
mock_sqlite.return_value = MagicMock()
|
||||
|
||||
from mem0.memory.main import Memory as MemoryClass
|
||||
memory = MemoryClass(MemoryConfig())
|
||||
|
||||
memory.add("foo", user_id="test_user", infer=False)
|
||||
|
||||
assert embedder.embed.call_count == 1
|
||||
mock_vector_store.insert.assert_called_once()
|
||||
|
||||
|
||||
@patch('mem0.utils.factory.EmbedderFactory.create')
|
||||
@patch('mem0.utils.factory.VectorStoreFactory.create')
|
||||
@patch('mem0.utils.factory.LlmFactory.create')
|
||||
@patch('mem0.memory.storage.SQLiteManager')
|
||||
def test_add_infer_true_caches_embedding_on_llm_rewrite(mock_sqlite, mock_llm_factory, mock_vector_factory, mock_embedder_factory):
|
||||
"""
|
||||
Regression test for issue #3723 (infer=True path): when the LLM rewrites a fact during the
|
||||
ADD action, the embedding should be computed once and cached, not computed again inside _create_memory.
|
||||
"""
|
||||
embedder = MagicMock()
|
||||
embedder.embed.return_value = [0.1, 0.2, 0.3]
|
||||
embedder.config = MagicMock(embedding_dims=3)
|
||||
mock_embedder_factory.return_value = embedder
|
||||
|
||||
mock_vector_store = MagicMock()
|
||||
mock_vector_store.search.return_value = []
|
||||
mock_vector_store.insert.return_value = None
|
||||
mock_vector_store.get.return_value = None
|
||||
telemetry_vector_store = MagicMock()
|
||||
mock_vector_factory.side_effect = [mock_vector_store, telemetry_vector_store]
|
||||
|
||||
# LLM extracts fact "User likes Python", then ADD action rewrites to "The user enjoys Python"
|
||||
mock_llm = MagicMock()
|
||||
mock_llm.generate_response.side_effect = [
|
||||
json.dumps({"facts": ["User likes Python"]}),
|
||||
json.dumps({"memory": [{"id": "0", "text": "The user enjoys Python", "event": "ADD", "old_memory": None}]}),
|
||||
]
|
||||
mock_llm_factory.return_value = mock_llm
|
||||
|
||||
mock_sqlite.return_value = MagicMock()
|
||||
|
||||
from mem0.memory.main import Memory as MemoryClass
|
||||
memory = MemoryClass(MemoryConfig())
|
||||
|
||||
memory.add("I like Python", user_id="test_user", infer=True)
|
||||
|
||||
# embed should be called exactly twice:
|
||||
# 1. For the extracted fact "User likes Python" (search)
|
||||
# 2. For the rewritten text "The user enjoys Python" (pre-cached before _create_memory)
|
||||
# It should NOT be called a 3rd time inside _create_memory
|
||||
assert embedder.embed.call_count == 2
|
||||
mock_vector_store.insert.assert_called_once()
|
||||
|
||||
Reference in New Issue
Block a user