eolas / app /retrieve.py
rahulraj1406's picture
v0.4.0-alpha: cross-encoder reranking + eval harness with CI gate
310c2f8
Raw
History Blame Contribute Delete
4.01 kB
"""
Retrieval — embed query, hit LanceDB, apply the RELEVANCE GATE (Safeguard 1).
The gate is the most important defense in this whole system. It runs BEFORE
the LLM is called: if the best retrieved passage scores below the threshold,
we never invoke the LLM, we just return "I don't know" plus a contact pointer.
That kills the worst failure mode (confident hallucinations on topics not in
the corpus) at zero LLM cost.
Score model: vectors are L2-normalised, so cosine similarity = dot product.
LanceDB returns _distance = 1 - cosine_similarity for cosine metric, so
similarity = 1 - _distance.
"""
from __future__ import annotations
from dataclasses import dataclass
from functools import lru_cache
import lancedb
from app.config import SETTINGS
from app.providers.local_bge import LocalBGE
# ── single shared embedder ────────────────────────────────────────────────────
@lru_cache(maxsize=1)
def get_embedder() -> LocalBGE:
return LocalBGE(SETTINGS.embed_model)
@lru_cache(maxsize=1)
def get_table():
db = lancedb.connect(str(SETTINGS.lancedb_path))
return db.open_table(SETTINGS.table_name)
# ── result shape ──────────────────────────────────────────────────────────────
@dataclass
class Passage:
text: str
source_url: str
title: str
heading_path: str
page: int | None
category: str
score: float
rerank_score: float | None = None
@dataclass
class RetrievalResult:
passages: list[Passage]
best_score: float
gate_passed: bool
# ── core ──────────────────────────────────────────────────────────────────────
def retrieve(
query: str,
top_k: int | None = None,
threshold: float | None = None,
) -> RetrievalResult:
"""Two-stage retrieval: dense top-N → cross-encoder rerank → top-K.
The relevance gate runs on the DENSE score of the best chunk (the bar for
"is this in our corpus at all"), not the rerank score. Reranking is a
re-sort, not a recall booster.
"""
top_k = top_k or SETTINGS.top_k
threshold = threshold if threshold is not None else SETTINGS.relevance_threshold
rerank_pool = max(SETTINGS.rerank_top_n, top_k)
embedder = get_embedder()
table = get_table()
vec = embedder.embed_query(query).tolist()
rows = (
table.search(vec)
.limit(rerank_pool)
.to_list()
)
passages: list[Passage] = []
for r in rows:
dist = float(r.get("_distance", 1.0))
# LanceDB cosine _distance is (1 - cos_sim) for L2-normalised vectors.
score = max(0.0, 1.0 - dist)
passages.append(Passage(
text=r["text"],
source_url=r["source_url"],
title=r["title"],
heading_path=r["heading_path"],
page=None if r.get("page", -1) == -1 else r["page"],
category=r["category"],
score=score,
))
best_dense = passages[0].score if passages else 0.0
gate_passed = best_dense >= threshold
# Rerank only when the gate passes; if it failed, we won't use these chunks.
if gate_passed and SETTINGS.rerank_enabled and len(passages) > 1:
from app.providers.reranker import rerank as _rerank
order = _rerank(query, [p.text for p in passages], top_k=top_k)
reranked: list[Passage] = []
for idx, rscore in order:
p = passages[idx]
p.rerank_score = rscore
reranked.append(p)
passages = reranked
else:
passages = passages[:top_k]
return RetrievalResult(
passages=passages,
best_score=best_dense,
gate_passed=gate_passed,
)