Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions config/chat/tests/test_agent_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,8 +89,10 @@ class FakeMemory:
def __init__(self):
self.added_messages = None
self.added_context = None
self.recall_rerank = None

def recall(self, *, query, context, limit):
def recall(self, *, query, context, limit, rerank=False):
self.recall_rerank = rerank
return [
MemorySearchResult(
id="memory-1",
Expand Down Expand Up @@ -162,6 +164,7 @@ def add(self, *, messages, context, prompt=None):

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

def test_memory_failures_do_not_abort_the_agent_turn(self):
class BrokenMemory:
def recall(self, *, query, context, limit):
def recall(self, *, query, context, limit, rerank=False):
return []

def add(self, *, messages, context, prompt=None):
Expand Down
3 changes: 3 additions & 0 deletions config/core/agent_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ def recall(
query: str,
context: MemoryContext,
limit: int = 5,
rerank: bool = False,
) -> list[MemorySearchResult]: ...

def add(
Expand All @@ -65,6 +66,7 @@ class AgentRuntimeConfig:
base_url: str | None = None
max_tokens: int = 8_000
max_rounds: int = MAX_ROUNDS
memory_rerank: bool = True

def __post_init__(self) -> None:
if not self.model:
Expand Down Expand Up @@ -285,6 +287,7 @@ def _system_with_recall(self, latest_user_query: str) -> str:
query=latest_user_query,
context=self.memory_context,
limit=5,
rerank=self.config.memory_rerank,
)
except Exception:
logger.warning(
Expand Down
73 changes: 73 additions & 0 deletions config/core/memory/bge_reranker.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
from __future__ import annotations

from collections.abc import Callable, Sequence
from dataclasses import replace
from typing import Protocol

from core.memory.vector_store import MemorySearchResult

DEFAULT_BGE_RERANKER_MODEL = "BAAI/bge-reranker-v2-m3"


class BGEScoringModel(Protocol):
def compute_score(
self,
sentence_pairs: list[list[str]],
*,
normalize: bool,
) -> Sequence[float] | float: ...


def _load_bge_model(model_name: str) -> BGEScoringModel:
from FlagEmbedding import FlagReranker # type: ignore[import-untyped]

return FlagReranker(model_name, use_fp16=False)


class BGEReranker:
"""Lazily load BGE and rerank Memory retrieval candidates."""

def __init__(
self,
model_name: str = DEFAULT_BGE_RERANKER_MODEL,
*,
model_factory: Callable[[str], BGEScoringModel] = _load_bge_model,
) -> None:
self.model_name = model_name
self._model_factory = model_factory
self._model: BGEScoringModel | None = None

def rerank(
self,
*,
query: str,
candidates: Sequence[MemorySearchResult],
limit: int,
) -> list[MemorySearchResult]:
if limit <= 0 or not candidates:
return []

raw_scores = self._get_model().compute_score(
[[query, candidate.data] for candidate in candidates],
normalize=True,
)
if isinstance(raw_scores, (float, int)):
scores = [float(raw_scores)]
else:
scores = [float(score) for score in raw_scores]
if len(scores) != len(candidates):
raise RuntimeError(
"BGE returned a score count that does not match candidates"
)

ranked = sorted(
zip(candidates, scores, strict=True),
key=lambda item: item[1],
reverse=True,
)
return [replace(candidate, score=score) for candidate, score in ranked[:limit]]

def _get_model(self) -> BGEScoringModel:
if self._model is None:
self._model = self._model_factory(self.model_name)
return self._model
2 changes: 2 additions & 0 deletions config/core/memory/composition.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from fastembed import SparseTextEmbedding
from qdrant_client import QdrantClient

from core.memory.bge_reranker import BGEReranker
from core.memory.config import AnthropicLLMConfig
from core.memory.embedder import (
MULTILINGUAL_E5_BASE_DIMENSION,
Expand Down Expand Up @@ -62,4 +63,5 @@ def build_memory(*, config: MemoryCompositionConfig) -> Memory:
extractor=extractor,
dense_encoder=dense_encoder,
vector_store=vector_store,
reranker=BGEReranker(),
)
99 changes: 99 additions & 0 deletions config/core/memory/flashrank_reranker.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,99 @@
from __future__ import annotations

from collections.abc import Callable, Mapping, Sequence
from dataclasses import replace
from importlib import import_module
from pathlib import Path
from typing import Any, Protocol, cast

from core.memory.vector_store import MemorySearchResult

DEFAULT_FLASHRANK_MODEL = "ms-marco-MultiBERT-L-12"
DEFAULT_FLASHRANK_CACHE_DIR = Path.home() / ".mini-code-agent" / "models" / "flashrank"


class FlashRankBackend(Protocol):
def __call__(
self,
*,
query: str,
passages: list[dict[str, object]],
) -> Sequence[Mapping[str, object]]: ...


def _load_flashrank_backend(
model_name: str,
cache_dir: Path,
) -> FlashRankBackend:
flashrank = import_module("flashrank")
ranker = flashrank.Ranker(model_name=model_name, cache_dir=str(cache_dir))

def rerank(
*,
query: str,
passages: list[dict[str, object]],
) -> Sequence[Mapping[str, object]]:
request = flashrank.RerankRequest(query=query, passages=passages)
return cast(Sequence[Mapping[str, object]], ranker.rerank(request))

return rerank


class FlashRankReranker:
"""Lazily load FlashRank and rerank Memory retrieval candidates."""

def __init__(
self,
model_name: str = DEFAULT_FLASHRANK_MODEL,
*,
cache_dir: Path = DEFAULT_FLASHRANK_CACHE_DIR,
backend_factory: Callable[[str, Path], FlashRankBackend] = (
_load_flashrank_backend
),
) -> None:
self.model_name = model_name
self.cache_dir = cache_dir
self._backend_factory = backend_factory
self._backend: FlashRankBackend | None = None

def rerank(
self,
*,
query: str,
candidates: Sequence[MemorySearchResult],
limit: int,
) -> list[MemorySearchResult]:
if limit <= 0 or not candidates:
return []

raw_results = self._get_backend()(
query=query,
passages=[
{"id": index, "text": candidate.data}
for index, candidate in enumerate(candidates)
],
)
if len(raw_results) != len(candidates):
raise RuntimeError(
"FlashRank returned a result count that does not match candidates"
)

try:
ranked = [
replace(
candidates[int(cast(Any, result["id"]))],
score=float(cast(Any, result["score"])),
)
for result in raw_results
]
except (KeyError, TypeError, ValueError, IndexError) as error:
raise RuntimeError("FlashRank returned an invalid result") from error
return ranked[:limit]

def _get_backend(self) -> FlashRankBackend:
if self._backend is None:
self._backend = self._backend_factory(
self.model_name,
self.cache_dir,
)
return self._backend
24 changes: 20 additions & 4 deletions config/core/memory/memory.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
MemoryExtractor,
MemoryMessage,
)
from core.memory.reranker import MemoryReranker
from core.memory.vector_store import MemorySearchResult, MemoryVectorStore


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


class Memory:
_RERANK_CANDIDATES_PER_SCOPE = 10

def __init__(
self,
*,
extractor: MemoryExtractor,
dense_encoder: DenseEncoder,
vector_store: MemoryVectorStore,
reranker: MemoryReranker | None = None,
) -> None:
self.extractor = extractor
self.dense_encoder = dense_encoder
self.vector_store = vector_store
self.reranker = reranker

def close(self) -> None:
"""Release resources owned by the configured vector-store adapter."""
Expand Down Expand Up @@ -174,20 +179,22 @@ def recall(
query: str,
context: MemoryContext,
limit: int = 5,
rerank: bool = False,
) -> builtins.list[MemorySearchResult]:
"""Recall User and current Space Memories for one agent Turn."""
scope_limit = self._RERANK_CANDIDATES_PER_SCOPE if rerank else limit
user_results = self._search_scope(
query=query,
filters={"user_id": context.user_id, "space_id": None},
limit=limit,
limit=scope_limit,
)
space_results = self._search_scope(
query=query,
filters={
"user_id": context.user_id,
"space_id": context.space_id,
},
limit=limit,
limit=scope_limit,
)

result_by_text: dict[str, MemorySearchResult] = {}
Expand All @@ -208,11 +215,20 @@ def recall(
score_by_text,
key=score_by_text.__getitem__,
reverse=True,
)[:limit]
return [
)[: scope_limit * 2]
candidates = [
replace(result_by_text[text], score=score_by_text[text])
for text in ranked_texts
]
if not rerank:
return candidates[:limit]
if self.reranker is None:
raise RuntimeError("rerank=True requires a configured MemoryReranker")
return self.reranker.rerank(
query=query,
candidates=candidates,
limit=limit,
)

def _reciprocal_rank_fusion(
self,
Expand Down
16 changes: 16 additions & 0 deletions config/core/memory/reranker.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,16 @@
from collections.abc import Sequence
from typing import Protocol

from core.memory.vector_store import MemorySearchResult


class MemoryReranker(Protocol):
"""Reorder retrieval candidates and return at most ``limit`` results."""

def rerank(
self,
*,
query: str,
candidates: Sequence[MemorySearchResult],
limit: int,
) -> list[MemorySearchResult]: ...
Loading