Open curriculum · verified editionContribute on GitHub
Lab20 min19 words

Lab 04 — Two-stage retrieval: cheap candidates and precise reranking

Localized walkthrough. Run the canonical public lab locally and preserve the result as evidence.

SourceImprove this page

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())