Skip to content

Commit 36f705d

Browse files
feat: neural reranker and NVIDIA nemotron-embed support
Add cross-encoder reranking stage using NVIDIA NIM API for improved retrieval accuracy. Add nemotron-embed model support with proper modality list handling in batch embeddings. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
1 parent e18aed4 commit 36f705d

6 files changed

Lines changed: 217 additions & 7 deletions

File tree

engram/benchmarks/longmemeval.py

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@
2727
LLMConfig,
2828
MemoryConfig,
2929
ProfileConfig,
30+
RerankConfig,
3031
SceneConfig,
3132
VectorStoreConfig,
3233
)
@@ -163,6 +164,8 @@ def build_memory(
163164
embedder_model: Optional[str] = None,
164165
full_potential: bool = True,
165166
defer_enrichment: bool = False,
167+
enable_rerank: bool = False,
168+
rerank_model: Optional[str] = None,
166169
) -> Memory:
167170
"""Build Engram Memory for LongMemEval. By default uses full potential (echo, categories, graph, scenes, profiles).
168171
@@ -174,13 +177,17 @@ def build_memory(
174177
"embedding_model_dims": embedding_dims,
175178
}
176179

177-
llm_cfg: Dict[str, Any] = {"max_tokens": 8192, "timeout": 300, "model": "meta/llama-3.3-70b-instruct"}
180+
llm_cfg: Dict[str, Any] = {"max_tokens": 16384, "timeout": 300, "model": "meta/llama-3.3-70b-instruct"}
178181
if llm_model:
179182
llm_cfg["model"] = llm_model
180183
embedder_cfg: Dict[str, Any] = {"embedding_dims": embedding_dims}
181184
if embedder_model:
182185
embedder_cfg["model"] = embedder_model
183186

187+
rerank_cfg = RerankConfig(enable_rerank=enable_rerank)
188+
if rerank_model:
189+
rerank_cfg = RerankConfig(enable_rerank=enable_rerank, model=rerank_model)
190+
184191
config = MemoryConfig(
185192
vector_store=VectorStoreConfig(provider=vector_store_provider, config=vector_cfg),
186193
llm=LLMConfig(provider=llm_provider, config=llm_cfg),
@@ -194,10 +201,11 @@ def build_memory(
194201
profile=ProfileConfig(use_llm_extraction=full_potential, enable_profiles=full_potential),
195202
enrichment=EnrichmentConfig(
196203
enable_unified=full_potential,
197-
max_batch_size=10,
204+
max_batch_size=5,
198205
defer_enrichment=defer_enrichment,
199206
),
200207
batch=BatchConfig(enable_batch=full_potential and not defer_enrichment, max_batch_size=50),
208+
rerank=rerank_cfg,
201209
)
202210
mem = Memory(config)
203211
# FullMemory features (categories, scenes, profiles) need FullSQLiteManager
@@ -267,6 +275,8 @@ def run_longmemeval(args: argparse.Namespace) -> Dict[str, Any]:
267275
embedder_model=args.embedder_model,
268276
full_potential=args.full_potential,
269277
defer_enrichment=use_deferred,
278+
enable_rerank=getattr(args, "enable_rerank", False),
279+
rerank_model=getattr(args, "rerank_model", None),
270280
)
271281

272282
hf_responder: Optional[HFResponder] = None
@@ -494,6 +504,8 @@ def parse_args() -> argparse.Namespace:
494504
parser.add_argument("--vector-store-provider", choices=["memory", "sqlite_vec"], default="memory")
495505
parser.add_argument("--history-db-path", default="/content/engram-longmemeval.db", help="SQLite db path.")
496506
parser.add_argument("--defer-enrichment", action="store_true", default=False, help="Use deferred enrichment (0 LLM calls at ingestion, batch enrich after).")
507+
parser.add_argument("--enable-rerank", action="store_true", default=False, help="Enable neural reranking (cross-encoder second stage on retrieved results).")
508+
parser.add_argument("--rerank-model", default=None, help="Reranker model override (default: nvidia/llama-3.2-nv-rerankqa-1b-v2).")
497509
args = parser.parse_args()
498510
args.full_potential = not args.minimal
499511
args.defer_enrichment = args.defer_enrichment

engram/configs/base.py

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -376,6 +376,16 @@ def _valid_priority(cls, v: str) -> str:
376376
return v
377377

378378

379+
class RerankConfig(BaseModel):
380+
"""Configuration for neural reranking (cross-encoder second stage)."""
381+
enable_rerank: bool = False
382+
provider: str = "nvidia" # Currently only nvidia supported
383+
model: str = "nvidia/llama-3.2-nv-rerankqa-1b-v2"
384+
api_key_env: str = "NVIDIA_API_KEY" # Env var name for API key
385+
top_n: int = 0 # Number of results to return after reranking (0 = return all, re-sorted)
386+
config: Dict[str, Any] = Field(default_factory=dict)
387+
388+
379389
class EnrichmentConfig(BaseModel):
380390
"""Configuration for unified enrichment (single LLM call for echo+category+entities+profiles)."""
381391
enable_unified: bool = False # Off by default for backward compat
@@ -481,6 +491,7 @@ class MemoryConfig(BaseModel):
481491
parallel: ParallelConfig = Field(default_factory=ParallelConfig)
482492
batch: BatchConfig = Field(default_factory=BatchConfig)
483493
enrichment: EnrichmentConfig = Field(default_factory=EnrichmentConfig)
494+
rerank: RerankConfig = Field(default_factory=RerankConfig)
484495
skill: SkillConfig = Field(default_factory=SkillConfig)
485496
task: TaskConfig = Field(default_factory=TaskConfig)
486497
metamemory: MetamemoryInlineConfig = Field(default_factory=MetamemoryInlineConfig)

engram/embeddings/nvidia.py

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -33,11 +33,20 @@ def __init__(self, config: Optional[dict] = None):
3333
self.client = OpenAI(base_url=base_url, api_key=api_key, timeout=timeout)
3434
self.model = self.config.get("model", "nvidia/nv-embed-v1")
3535

36-
def _extra_body(self, memory_action: Optional[str] = None) -> dict:
37-
"""Build extra_body for E5/embedqa models."""
36+
def _extra_body(self, memory_action: Optional[str] = None, count: int = 1) -> dict:
37+
"""Build extra_body for models that need input_type differentiation.
38+
39+
Args:
40+
memory_action: The action type (search, forget, etc.)
41+
count: Number of texts in the batch. nemotron-embed requires
42+
modality list length to match input length.
43+
"""
3844
if "e5" in self.model or "embedqa" in self.model:
3945
input_type = "query" if memory_action in ("search", "forget") else "passage"
4046
return {"input_type": input_type, "truncate": "END"}
47+
if "nemotron-embed" in self.model:
48+
input_type = "query" if memory_action in ("search", "forget") else "passage"
49+
return {"modality": ["text"] * count, "input_type": input_type, "truncate": "END"}
4150
return {}
4251

4352
def _truncate_if_needed(self, text: str) -> str:
@@ -89,7 +98,7 @@ def embed_batch(
8998
if len(texts) == 1:
9099
return [self.embed(texts[0], memory_action=memory_action)]
91100
try:
92-
extra_body = self._extra_body(memory_action)
101+
extra_body = self._extra_body(memory_action, count=len(texts))
93102
response = self.client.embeddings.create(
94103
input=texts,
95104
model=self.model,

engram/memory/main.py

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -254,6 +254,8 @@ def __init__(self, config: Optional[MemoryConfig] = None, preset: Optional[str]
254254
self._profile_processor: Optional[ProfileProcessor] = None
255255
self._task_manager: Optional[Any] = None
256256
self._project_manager: Optional[Any] = None
257+
# Neural reranker (lazy init)
258+
self._reranker: Optional[Any] = None
257259
# Trajectory recording and skill mining
258260
self._trajectory_store: Optional[Any] = None
259261
self._skill_miner: Optional[Any] = None
@@ -327,6 +329,20 @@ def skill_miner(self):
327329
)
328330
return self._skill_miner
329331

332+
@property
333+
def reranker(self):
334+
"""Lazy-initialized neural reranker (only if enabled in config)."""
335+
rerank_cfg = getattr(self.config, "rerank", None)
336+
if self._reranker is None and rerank_cfg and rerank_cfg.enable_rerank:
337+
from engram.retrieval.reranker import create_reranker
338+
self._reranker = create_reranker({
339+
"provider": rerank_cfg.provider,
340+
"model": rerank_cfg.model,
341+
"api_key_env": rerank_cfg.api_key_env,
342+
**rerank_cfg.config,
343+
})
344+
return self._reranker
345+
330346
def start_trajectory(
331347
self,
332348
task_description: str,
@@ -2367,6 +2383,28 @@ def search(
23672383

23682384
results.sort(key=lambda x: x["composite_score"], reverse=True)
23692385

2386+
# Neural reranking: cross-encoder second stage on top candidates
2387+
rerank_cfg = getattr(self.config, "rerank", None)
2388+
if rerank and self.reranker and results:
2389+
try:
2390+
passages = [r.get("memory", "") for r in results[:limit]]
2391+
reranked = self.reranker.rerank(
2392+
query=query,
2393+
passages=passages,
2394+
top_n=rerank_cfg.top_n if rerank_cfg and rerank_cfg.top_n > 0 else 0,
2395+
)
2396+
# Re-order results by reranker logits
2397+
idx_to_logit = {r["index"]: r["logit"] for r in reranked}
2398+
for i, result in enumerate(results[:limit]):
2399+
result["rerank_logit"] = idx_to_logit.get(i, float("-inf"))
2400+
results[:limit] = sorted(
2401+
results[:limit],
2402+
key=lambda x: x.get("rerank_logit", float("-inf")),
2403+
reverse=True,
2404+
)
2405+
except Exception as e:
2406+
logger.warning("Reranking failed, using composite_score order: %s", e)
2407+
23702408
# Metamemory: auto-log knowledge gap when search returns no results
23712409
if not results and self.config.metamemory.auto_log_gaps:
23722410
try:

engram/retrieval/__init__.py

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,10 @@
11
"""Engram v2 retrieval components."""
22

3-
from engram.retrieval.dual_search import DualSearchEngine
3+
try:
4+
from engram.retrieval.dual_search import DualSearchEngine
5+
except ImportError:
6+
DualSearchEngine = None
47

5-
__all__ = ["DualSearchEngine"]
8+
from engram.retrieval.reranker import NvidiaReranker, create_reranker
9+
10+
__all__ = ["DualSearchEngine", "NvidiaReranker", "create_reranker"]

engram/retrieval/reranker.py

Lines changed: 135 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,135 @@
1+
"""Neural reranker for second-stage retrieval refinement.
2+
3+
Uses a cross-encoder model to re-score (query, passage) pairs with full
4+
attention, producing much more accurate relevance scores than embedding
5+
cosine similarity alone.
6+
"""
7+
8+
import logging
9+
import os
10+
import time
11+
from typing import Any, Dict, List, Optional
12+
13+
import requests
14+
15+
logger = logging.getLogger(__name__)
16+
17+
18+
class NvidiaReranker:
19+
"""NVIDIA NIM reranker using the /reranking endpoint."""
20+
21+
_DEFAULT_URL = (
22+
"https://ai.api.nvidia.com/v1/retrieval/"
23+
"nvidia/llama-3_2-nv-rerankqa-1b-v2/reranking"
24+
)
25+
26+
def __init__(self, config: Optional[Dict[str, Any]] = None):
27+
config = config or {}
28+
self.model = config.get("model", "nvidia/llama-3.2-nv-rerankqa-1b-v2")
29+
api_key_env = config.get("api_key_env", "NVIDIA_API_KEY")
30+
self.api_key = config.get("api_key") or os.getenv(api_key_env)
31+
if not self.api_key:
32+
raise ValueError(
33+
f"NVIDIA API key required for reranker. Set config['api_key'] or {api_key_env} env var."
34+
)
35+
# Build URL from model name: replace / with _ and dots with _
36+
# e.g. nvidia/llama-3.2-nv-rerankqa-1b-v2 -> nvidia/llama-3_2-nv-rerankqa-1b-v2
37+
model_path = self.model.replace(".", "_")
38+
self.url = config.get(
39+
"url",
40+
f"https://ai.api.nvidia.com/v1/retrieval/{model_path}/reranking",
41+
)
42+
self.timeout = config.get("timeout", 30)
43+
self.max_retries = config.get("max_retries", 2)
44+
45+
def rerank(
46+
self,
47+
query: str,
48+
passages: List[str],
49+
top_n: int = 0,
50+
) -> List[Dict[str, Any]]:
51+
"""Rerank passages against a query.
52+
53+
Args:
54+
query: The search query.
55+
passages: List of passage texts to rerank.
56+
top_n: Number of top results to return (0 = return all, re-sorted).
57+
58+
Returns:
59+
List of dicts with keys: index (original position), logit, text.
60+
Sorted by logit descending.
61+
"""
62+
if not passages:
63+
return []
64+
if len(passages) == 1:
65+
return [{"index": 0, "logit": 0.0, "text": passages[0]}]
66+
67+
payload = {
68+
"model": self.model,
69+
"query": {"text": query},
70+
"passages": [{"text": p} for p in passages],
71+
}
72+
if top_n > 0:
73+
payload["top_n"] = top_n
74+
75+
headers = {
76+
"Authorization": f"Bearer {self.api_key}",
77+
"Content-Type": "application/json",
78+
"Accept": "application/json",
79+
}
80+
81+
last_exc = None
82+
for attempt in range(self.max_retries + 1):
83+
try:
84+
t0 = time.monotonic()
85+
resp = requests.post(
86+
self.url,
87+
json=payload,
88+
headers=headers,
89+
timeout=self.timeout,
90+
)
91+
elapsed_ms = (time.monotonic() - t0) * 1000
92+
resp.raise_for_status()
93+
data = resp.json()
94+
95+
rankings = data.get("rankings", [])
96+
results = []
97+
for r in rankings:
98+
idx = r.get("index", 0)
99+
results.append({
100+
"index": idx,
101+
"logit": r.get("logit", 0.0),
102+
"text": passages[idx] if idx < len(passages) else "",
103+
})
104+
results.sort(key=lambda x: x["logit"], reverse=True)
105+
logger.debug(
106+
"Reranked %d passages in %.0fms (top logit=%.2f)",
107+
len(passages), elapsed_ms,
108+
results[0]["logit"] if results else 0.0,
109+
)
110+
return results
111+
112+
except Exception as exc:
113+
last_exc = exc
114+
if attempt < self.max_retries:
115+
delay = min(2 ** attempt, 4)
116+
logger.warning(
117+
"Reranker retry %d/%d after %ss: %s",
118+
attempt + 1, self.max_retries, delay, exc,
119+
)
120+
time.sleep(delay)
121+
else:
122+
logger.error("Reranker failed after %d attempts: %s", self.max_retries + 1, exc)
123+
124+
raise RuntimeError(f"Reranker failed: {last_exc}") from last_exc
125+
126+
127+
def create_reranker(config: Optional[Dict[str, Any]] = None) -> Optional[NvidiaReranker]:
128+
"""Factory: create a reranker from config, or return None if disabled."""
129+
if not config:
130+
return None
131+
provider = config.get("provider", "nvidia")
132+
if provider == "nvidia":
133+
return NvidiaReranker(config)
134+
logger.warning("Unknown reranker provider: %s", provider)
135+
return None

0 commit comments

Comments
 (0)