Skip to content

Commit d20705d

Browse files
committed
fix: keep query embedding for fine search
1 parent 96a1dd6 commit d20705d

2 files changed

Lines changed: 38 additions & 5 deletions

File tree

src/memos/memories/textual/tree_text_memory/retrieve/searcher.py

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -320,10 +320,13 @@ def _parse_task(
320320
)
321321

322322
query = parsed_goal.rephrased_query or query
323-
# if goal has extra memories, embed them too
324-
if parsed_goal.memories:
325-
embed_texts = list(dict.fromkeys([query, *parsed_goal.memories]))
326-
query_embedding = self.embedder.embed(embed_texts)
323+
embed_texts = [
324+
text.strip()
325+
for text in [query, *parsed_goal.memories]
326+
if isinstance(text, str) and text.strip()
327+
]
328+
if embed_texts:
329+
query_embedding = self.embedder.embed(list(dict.fromkeys(embed_texts)))
327330
return parsed_goal, query_embedding, context, query
328331

329332
@timed

tests/memories/textual/test_tree_searcher.py

Lines changed: 31 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import pytest
44

55
from memos.memories.textual.item import TextualMemoryItem, TreeNodeTextualMemoryMetadata
6+
from memos.memories.textual.tree_text_memory.retrieve.retrieval_mid_structs import ParsedTaskGoal
67
from memos.memories.textual.tree_text_memory.retrieve.searcher import Searcher
78
from memos.reranker.base import BaseReranker
89

@@ -24,13 +25,14 @@ def mock_searcher():
2425
return s
2526

2627

27-
def make_item(content: str, score: float):
28+
def make_item(content: str, score: float, memory_type: str = "WorkingMemory"):
2829
# Simulate a TextualMemoryItem with usage list for update test
2930
return (
3031
TextualMemoryItem(
3132
memory=content,
3233
metadata=TreeNodeTextualMemoryMetadata(
3334
embedding=[0.1] * 5,
35+
memory_type=memory_type,
3436
usage=[],
3537
),
3638
),
@@ -104,6 +106,34 @@ def test_searcher_fine_mode_triggers_reasoner(mock_searcher):
104106
assert len(result) == 1
105107

106108

109+
def test_fine_search_embeds_query_when_parser_returns_no_memory_expansions(mock_searcher):
110+
query = "我喜欢什么"
111+
user_memory = make_item("我喜欢草莓", 0.9, memory_type="UserMemory")[0]
112+
113+
mock_searcher.task_goal_parser.parse.return_value = ParsedTaskGoal(
114+
keys=["喜欢", "偏好", "兴趣"],
115+
tags=["personal preference", "user interest", "taste"],
116+
memories=[],
117+
rephrased_query="",
118+
)
119+
mock_searcher.embedder.embed.return_value = [[0.1] * 5]
120+
121+
def retrieve_side_effect(*args, **kwargs):
122+
if kwargs.get("memory_scope") == "UserMemory":
123+
assert kwargs["query_embedding"] == [[0.1] * 5]
124+
return [user_memory]
125+
return []
126+
127+
mock_searcher.graph_retriever.retrieve.side_effect = retrieve_side_effect
128+
mock_searcher.reranker.rerank.return_value = [(user_memory, 0.9)]
129+
130+
result = mock_searcher.search(query=query, top_k=1, mode="fine", memory_type="UserMemory")
131+
132+
mock_searcher.embedder.embed.assert_called_once_with([query])
133+
assert len(result) == 1
134+
assert result[0].memory == "我喜欢草莓"
135+
136+
107137
def test_searcher_respects_memory_type(mock_searcher):
108138
parsed_goal = MagicMock()
109139
parsed_goal.memories = ["Something"]

0 commit comments

Comments
 (0)