Lab 04 — Two-stage retrieval: cheap candidates and precise reranking
Localized walkthrough. Run the canonical public lab locally and preserve the result as evidence.
Localized walkthrough. Run the canonical public lab locally and preserve the result as evidence.
Run it#
python modulo-03-rag/labs/04_reranking.py
Localized source walkthrough#
"""Lab 04 — Two-stage retrieval: cheap candidates and precise reranking.
The normal mode uses embeddings for retrieval and a multilingual cross-encoder for reranking.
``--lexical`` avoids downloads: TF-IDF retrieves and a deterministic function based on coverage,
title, and density reranks. This mode is a smoke test, not a substitute for the cross-encoder.
Execution:
python modulo-03-rag/labs/04_reranking.py --lexical
python modulo-03-rag/labs/04_reranking.py --limit 20
"""
from __future__ import annotations
import argparse
import json
import statistics
from dataclasses import dataclass
from _rag_common import (
DATA_DIR,
SearchHit,
SearchIndex,
heading_chunks,
load_documents,
tokenize,
)
from rich.console import Console
from rich.table import Table
DEFAULT_RERANKER = "cross-encoder/mmarco-mMiniLMv2-L12-H384-v1"
console = Console()
@dataclass(frozen=True)
class RankedHit:
hit: SearchHit
rerank_score: float
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"question",
nargs="?",
default="¿Qué debo hacer primero si se filtra una clave API?",
)
parser.add_argument("--candidate-k", type=int, default=10)
parser.add_argument("--final-k", type=int, default=4)
parser.add_argument("--limit", type=int, default=30)
parser.add_argument("--lexical", action="store_true")
parser.add_argument("--model", default=DEFAULT_RERANKER)
return parser.parse_args()
def lexical_rerank(question: str, hits: list[SearchHit]) -> list[RankedHit]:
"""Explainable fallback for CI: term coverage and title priority."""
query_terms = set(tokenize(question))
ranked = []
for hit in hits:
body_terms = tokenize(hit.chunk.text)
title_terms = set(tokenize(hit.chunk.title))
overlap = query_terms & set(body_terms)
coverage = len(overlap) / max(1, len(query_terms))
density = sum(term in overlap for term in body_terms) / max(1, len(body_terms))
title_bonus = len(query_terms & title_terms) / max(1, len(query_terms))
score = 0.65 * coverage + 0.20 * title_bonus + 0.10 * density + 0.05 * hit.score
ranked.append(RankedHit(hit=hit, rerank_score=score))
return sorted(ranked, key=lambda item: item.rerank_score, reverse=True)
class CrossEncoderReranker:
def __init__(self, model_name: str) -> None:
from sentence_transformers import CrossEncoder
self.model = CrossEncoder(model_name)
def rerank(self, question: str, hits: list[SearchHit]) -> list[RankedHit]:
pairs = [
(question, f"{hit.chunk.title}\n{hit.chunk.text}")
for hit in hits
]
scores = self.model.predict(pairs, show_progress_bar=False)
ranked = [
RankedHit(hit=hit, rerank_score=float(score))
for hit, score in zip(hits, scores, strict=True)
]
return sorted(ranked, key=lambda item: item.rerank_score, reverse=True)
def doc_metrics(retrieved: list[str], expected: set[str]) -> tuple[float, float]:
recall = len(set(retrieved) & expected) / max(1, len(expected))
first = next(
(rank for rank, doc_id in enumerate(retrieved, start=1) if doc_id in expected),
None,
)
return recall, 1 / first if first else 0.0
def evaluate(
index: SearchIndex,
cases: list[dict],
*,
candidate_k: int,
final_k: int,
lexical: bool,
reranker: CrossEncoderReranker | None,
) -> dict[str, float]:
baseline_recalls: list[float] = []
baseline_rrs: list[float] = []
reranked_recalls: list[float] = []
reranked_rrs: list[float] = []
for case in cases:
candidates = index.search(case["question"], top_k=candidate_k)
ordered = (
lexical_rerank(case["question"], candidates)
if lexical
else reranker.rerank(case["question"], candidates)
)
expected = set(case["relevant_doc_ids"])
baseline_ids = [hit.chunk.doc_id for hit in candidates[:final_k]]
reranked_ids = [item.hit.chunk.doc_id for item in ordered[:final_k]]
baseline_recall, baseline_rr = doc_metrics(baseline_ids, expected)
reranked_recall, reranked_rr = doc_metrics(reranked_ids, expected)
baseline_recalls.append(baseline_recall)
baseline_rrs.append(baseline_rr)
reranked_recalls.append(reranked_recall)
reranked_rrs.append(reranked_rr)
return {
"baseline_recall": statistics.mean(baseline_recalls),
"baseline_mrr": statistics.mean(baseline_rrs),
"reranked_recall": statistics.mean(reranked_recalls),
"reranked_mrr": statistics.mean(reranked_rrs),
}
def validate_args(args: argparse.Namespace) -> None:
if args.final_k <= 0 or args.candidate_k < args.final_k:
raise ValueError("se requiere candidate-k >= final-k > 0")
if args.limit <= 0:
raise ValueError("limit debe ser > 0")
def main() -> int:
args = parse_args()
validate_args(args)
chunks = heading_chunks(load_documents())
index = SearchIndex(chunks, lexical=args.lexical)
reranker = None if args.lexical else CrossEncoderReranker(args.model)
candidates = index.search(args.question, top_k=args.candidate_k)
ranked = (
lexical_rerank(args.question, candidates)
if args.lexical
else reranker.rerank(args.question, candidates)
)
before = {hit.chunk.chunk_id: rank for rank, hit in enumerate(candidates, start=1)}
table = Table(title="Retrieval → reranking")
table.add_column("final")
table.add_column("inicial")
table.add_column("chunk")
table.add_column("retrieval", justify="right")
table.add_column("reranker", justify="right")
for rank, item in enumerate(ranked[: args.final_k], start=1):
table.add_row(
str(rank),
str(before[item.hit.chunk.chunk_id]),
item.hit.chunk.chunk_id,
f"{item.hit.score:.3f}",
f"{item.rerank_score:.3f}",
)
console.print(table)
with (DATA_DIR / "eval_dataset.json").open(encoding="utf-8") as handle:
cases = json.load(handle)[: args.limit]
metrics = evaluate(
index,
cases,
candidate_k=args.candidate_k,
final_k=args.final_k,
lexical=args.lexical,
reranker=reranker,
)
summary = Table(title=f"Evaluación · {len(cases)} casos · final@{args.final_k}")
summary.add_column("pipeline")
summary.add_column("doc recall", justify="right")
summary.add_column("MRR", justify="right")
summary.add_row(
"retrieval",
f"{metrics['baseline_recall']:.1%}",
f"{metrics['baseline_mrr']:.3f}",
)
summary.add_row(
"retrieval + reranker",
f"{metrics['reranked_recall']:.1%}",
f"{metrics['reranked_mrr']:.3f}",
)
console.print(summary)
if args.lexical:
console.print("[dim]Modo offline: el reranker es una heurística auditable.[/dim]")
return 0
if __name__ == "__main__":
raise SystemExit(main())