Skip to content

Commit 7510d8c

Browse files
authored
Merge pull request #1 from JohnnyYwQ/feat/memory-eval
feat: add memory reranking and retrieval evals
2 parents e29a807 + 627c248 commit 7510d8c

17 files changed

Lines changed: 2931 additions & 48 deletions

config/chat/tests/test_agent_runtime.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -89,8 +89,10 @@ class FakeMemory:
8989
def __init__(self):
9090
self.added_messages = None
9191
self.added_context = None
92+
self.recall_rerank = None
9293

93-
def recall(self, *, query, context, limit):
94+
def recall(self, *, query, context, limit, rerank=False):
95+
self.recall_rerank = rerank
9496
return [
9597
MemorySearchResult(
9698
id="memory-1",
@@ -162,6 +164,7 @@ def add(self, *, messages, context, prompt=None):
162164

163165
first_request = create.call_args_list[0].kwargs
164166
self.assertIn("[user] User prefers concise answers.", first_request["system"])
167+
self.assertIs(memory.recall_rerank, True)
165168
remember_tool = next(
166169
tool for tool in first_request["tools"] if tool["name"] == "remember"
167170
)
@@ -188,7 +191,7 @@ def add(self, *, messages, context, prompt=None):
188191

189192
def test_memory_failures_do_not_abort_the_agent_turn(self):
190193
class BrokenMemory:
191-
def recall(self, *, query, context, limit):
194+
def recall(self, *, query, context, limit, rerank=False):
192195
return []
193196

194197
def add(self, *, messages, context, prompt=None):

config/core/agent_runtime.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@ def recall(
4242
query: str,
4343
context: MemoryContext,
4444
limit: int = 5,
45+
rerank: bool = False,
4546
) -> list[MemorySearchResult]: ...
4647

4748
def add(
@@ -65,6 +66,7 @@ class AgentRuntimeConfig:
6566
base_url: str | None = None
6667
max_tokens: int = 8_000
6768
max_rounds: int = MAX_ROUNDS
69+
memory_rerank: bool = True
6870

6971
def __post_init__(self) -> None:
7072
if not self.model:
@@ -285,6 +287,7 @@ def _system_with_recall(self, latest_user_query: str) -> str:
285287
query=latest_user_query,
286288
context=self.memory_context,
287289
limit=5,
290+
rerank=self.config.memory_rerank,
288291
)
289292
except Exception:
290293
logger.warning(

config/core/memory/bge_reranker.py

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,73 @@
1+
from __future__ import annotations
2+
3+
from collections.abc import Callable, Sequence
4+
from dataclasses import replace
5+
from typing import Protocol
6+
7+
from core.memory.vector_store import MemorySearchResult
8+
9+
DEFAULT_BGE_RERANKER_MODEL = "BAAI/bge-reranker-v2-m3"
10+
11+
12+
class BGEScoringModel(Protocol):
13+
def compute_score(
14+
self,
15+
sentence_pairs: list[list[str]],
16+
*,
17+
normalize: bool,
18+
) -> Sequence[float] | float: ...
19+
20+
21+
def _load_bge_model(model_name: str) -> BGEScoringModel:
22+
from FlagEmbedding import FlagReranker # type: ignore[import-untyped]
23+
24+
return FlagReranker(model_name, use_fp16=False)
25+
26+
27+
class BGEReranker:
28+
"""Lazily load BGE and rerank Memory retrieval candidates."""
29+
30+
def __init__(
31+
self,
32+
model_name: str = DEFAULT_BGE_RERANKER_MODEL,
33+
*,
34+
model_factory: Callable[[str], BGEScoringModel] = _load_bge_model,
35+
) -> None:
36+
self.model_name = model_name
37+
self._model_factory = model_factory
38+
self._model: BGEScoringModel | None = None
39+
40+
def rerank(
41+
self,
42+
*,
43+
query: str,
44+
candidates: Sequence[MemorySearchResult],
45+
limit: int,
46+
) -> list[MemorySearchResult]:
47+
if limit <= 0 or not candidates:
48+
return []
49+
50+
raw_scores = self._get_model().compute_score(
51+
[[query, candidate.data] for candidate in candidates],
52+
normalize=True,
53+
)
54+
if isinstance(raw_scores, (float, int)):
55+
scores = [float(raw_scores)]
56+
else:
57+
scores = [float(score) for score in raw_scores]
58+
if len(scores) != len(candidates):
59+
raise RuntimeError(
60+
"BGE returned a score count that does not match candidates"
61+
)
62+
63+
ranked = sorted(
64+
zip(candidates, scores, strict=True),
65+
key=lambda item: item[1],
66+
reverse=True,
67+
)
68+
return [replace(candidate, score=score) for candidate, score in ranked[:limit]]
69+
70+
def _get_model(self) -> BGEScoringModel:
71+
if self._model is None:
72+
self._model = self._model_factory(self.model_name)
73+
return self._model

config/core/memory/composition.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from fastembed import SparseTextEmbedding
88
from qdrant_client import QdrantClient
99

10+
from core.memory.bge_reranker import BGEReranker
1011
from core.memory.config import AnthropicLLMConfig
1112
from core.memory.embedder import (
1213
MULTILINGUAL_E5_BASE_DIMENSION,
@@ -62,4 +63,5 @@ def build_memory(*, config: MemoryCompositionConfig) -> Memory:
6263
extractor=extractor,
6364
dense_encoder=dense_encoder,
6465
vector_store=vector_store,
66+
reranker=BGEReranker(),
6567
)
Lines changed: 99 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,99 @@
1+
from __future__ import annotations
2+
3+
from collections.abc import Callable, Mapping, Sequence
4+
from dataclasses import replace
5+
from importlib import import_module
6+
from pathlib import Path
7+
from typing import Any, Protocol, cast
8+
9+
from core.memory.vector_store import MemorySearchResult
10+
11+
DEFAULT_FLASHRANK_MODEL = "ms-marco-MultiBERT-L-12"
12+
DEFAULT_FLASHRANK_CACHE_DIR = Path.home() / ".mini-code-agent" / "models" / "flashrank"
13+
14+
15+
class FlashRankBackend(Protocol):
16+
def __call__(
17+
self,
18+
*,
19+
query: str,
20+
passages: list[dict[str, object]],
21+
) -> Sequence[Mapping[str, object]]: ...
22+
23+
24+
def _load_flashrank_backend(
25+
model_name: str,
26+
cache_dir: Path,
27+
) -> FlashRankBackend:
28+
flashrank = import_module("flashrank")
29+
ranker = flashrank.Ranker(model_name=model_name, cache_dir=str(cache_dir))
30+
31+
def rerank(
32+
*,
33+
query: str,
34+
passages: list[dict[str, object]],
35+
) -> Sequence[Mapping[str, object]]:
36+
request = flashrank.RerankRequest(query=query, passages=passages)
37+
return cast(Sequence[Mapping[str, object]], ranker.rerank(request))
38+
39+
return rerank
40+
41+
42+
class FlashRankReranker:
43+
"""Lazily load FlashRank and rerank Memory retrieval candidates."""
44+
45+
def __init__(
46+
self,
47+
model_name: str = DEFAULT_FLASHRANK_MODEL,
48+
*,
49+
cache_dir: Path = DEFAULT_FLASHRANK_CACHE_DIR,
50+
backend_factory: Callable[[str, Path], FlashRankBackend] = (
51+
_load_flashrank_backend
52+
),
53+
) -> None:
54+
self.model_name = model_name
55+
self.cache_dir = cache_dir
56+
self._backend_factory = backend_factory
57+
self._backend: FlashRankBackend | None = None
58+
59+
def rerank(
60+
self,
61+
*,
62+
query: str,
63+
candidates: Sequence[MemorySearchResult],
64+
limit: int,
65+
) -> list[MemorySearchResult]:
66+
if limit <= 0 or not candidates:
67+
return []
68+
69+
raw_results = self._get_backend()(
70+
query=query,
71+
passages=[
72+
{"id": index, "text": candidate.data}
73+
for index, candidate in enumerate(candidates)
74+
],
75+
)
76+
if len(raw_results) != len(candidates):
77+
raise RuntimeError(
78+
"FlashRank returned a result count that does not match candidates"
79+
)
80+
81+
try:
82+
ranked = [
83+
replace(
84+
candidates[int(cast(Any, result["id"]))],
85+
score=float(cast(Any, result["score"])),
86+
)
87+
for result in raw_results
88+
]
89+
except (KeyError, TypeError, ValueError, IndexError) as error:
90+
raise RuntimeError("FlashRank returned an invalid result") from error
91+
return ranked[:limit]
92+
93+
def _get_backend(self) -> FlashRankBackend:
94+
if self._backend is None:
95+
self._backend = self._backend_factory(
96+
self.model_name,
97+
self.cache_dir,
98+
)
99+
return self._backend

config/core/memory/memory.py

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
MemoryExtractor,
1313
MemoryMessage,
1414
)
15+
from core.memory.reranker import MemoryReranker
1516
from core.memory.vector_store import MemorySearchResult, MemoryVectorStore
1617

1718

@@ -37,16 +38,20 @@ def __post_init__(self) -> None:
3738

3839

3940
class Memory:
41+
_RERANK_CANDIDATES_PER_SCOPE = 10
42+
4043
def __init__(
4144
self,
4245
*,
4346
extractor: MemoryExtractor,
4447
dense_encoder: DenseEncoder,
4548
vector_store: MemoryVectorStore,
49+
reranker: MemoryReranker | None = None,
4650
) -> None:
4751
self.extractor = extractor
4852
self.dense_encoder = dense_encoder
4953
self.vector_store = vector_store
54+
self.reranker = reranker
5055

5156
def close(self) -> None:
5257
"""Release resources owned by the configured vector-store adapter."""
@@ -174,20 +179,22 @@ def recall(
174179
query: str,
175180
context: MemoryContext,
176181
limit: int = 5,
182+
rerank: bool = False,
177183
) -> builtins.list[MemorySearchResult]:
178184
"""Recall User and current Space Memories for one agent Turn."""
185+
scope_limit = self._RERANK_CANDIDATES_PER_SCOPE if rerank else limit
179186
user_results = self._search_scope(
180187
query=query,
181188
filters={"user_id": context.user_id, "space_id": None},
182-
limit=limit,
189+
limit=scope_limit,
183190
)
184191
space_results = self._search_scope(
185192
query=query,
186193
filters={
187194
"user_id": context.user_id,
188195
"space_id": context.space_id,
189196
},
190-
limit=limit,
197+
limit=scope_limit,
191198
)
192199

193200
result_by_text: dict[str, MemorySearchResult] = {}
@@ -208,11 +215,20 @@ def recall(
208215
score_by_text,
209216
key=score_by_text.__getitem__,
210217
reverse=True,
211-
)[:limit]
212-
return [
218+
)[: scope_limit * 2]
219+
candidates = [
213220
replace(result_by_text[text], score=score_by_text[text])
214221
for text in ranked_texts
215222
]
223+
if not rerank:
224+
return candidates[:limit]
225+
if self.reranker is None:
226+
raise RuntimeError("rerank=True requires a configured MemoryReranker")
227+
return self.reranker.rerank(
228+
query=query,
229+
candidates=candidates,
230+
limit=limit,
231+
)
216232

217233
def _reciprocal_rank_fusion(
218234
self,

config/core/memory/reranker.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,16 @@
1+
from collections.abc import Sequence
2+
from typing import Protocol
3+
4+
from core.memory.vector_store import MemorySearchResult
5+
6+
7+
class MemoryReranker(Protocol):
8+
"""Reorder retrieval candidates and return at most ``limit`` results."""
9+
10+
def rerank(
11+
self,
12+
*,
13+
query: str,
14+
candidates: Sequence[MemorySearchResult],
15+
limit: int,
16+
) -> list[MemorySearchResult]: ...

0 commit comments

Comments
 (0)