RAG & LLMs / 6. CROSS-ENCODER RERANKING

Stage 6: Cross-Encoder Reranking

The accuracy booster — deep pairwise (query, chunk) scoring


EXPLANATION

Retrieval gives you top-10 chunks fast but approximately. Cross-encoder reranking scores each (query, chunk) pair through a full BERT encoder together — so the model sees how query and chunk relate to each other at every attention layer.

This is far more accurate than cosine similarity but expensive (can't cache vectors). That's why you do two stages:

1. Bi-encoder retrieves top-10 (fast, approximate)
2. Cross-encoder re-scores all 10 (accurate, exact)
3. Keep top-3, discard the rest
4. Pass those 3 to the LLM

This keeps context small and quality high.

DATA FLOW

10 chunks from hybrid retrieval
          ↓
  Cross-Encoder: [CLS] query [SEP] chunk [SEP] → BERT → score

  chunk_1  score: 0.95  ██████████  ← highly relevant
  chunk_3  score: 0.89  █████████   ← very relevant
  chunk_4  score: 0.78  ████████    ← relevant
  chunk_2  score: 0.12  █           ← not relevant
  chunk_5  score: 0.03  ░           ← irrelevant
          ↓
  Pass top-3 to LLM (smaller context = better answer)

CODE

PYTHON
1from sentence_transformers import CrossEncoder
2
3# ── Load cross-encoder (BERT fine-tuned for passage ranking) ──────
4cross_encoder = CrossEncoder(
5 "cross-encoder/ms-marco-MiniLM-L-6-v2", # fast + good
6 # "cross-encoder/ms-marco-electra-base" # slower + better
7 max_length=512,
8)
9
10def rerank(query: str, docs: list, top_n: int = 3) -> list:
11 """
12 Rerank retrieved documents using cross-encoder.
13 Returns top_n most relevant documents.
14 """
15 if not docs:
16 return []
17
18 # Create (query, passage) pairs for the cross-encoder
19 pairs = [(query, doc.page_content) for doc in docs]
20
21 # Score all pairs — reads query + chunk TOGETHER through BERT
22 scores = cross_encoder.predict(pairs)
23
24 # Sort by score descending
25 ranked = sorted(zip(scores, docs), key=lambda x: x[0], reverse=True)
26
27 print(f"\nReranking {len(docs)} top {top_n}:")
28 for i, (score, doc) in enumerate(ranked[:top_n]):
29 print(f" [{i+1}] {score:.4f} {doc.page_content[:70]}...")
30
31 return [doc for _, doc in ranked[:top_n]]
32
33# ── Full retrieval + reranking ────────────────────────────────────
34query = "How does attention mechanism work in transformers?"
35
36retrieved = hybrid_retriever.invoke(query) # 10 chunks, fast
37reranked = rerank(query, retrieved, top_n=3) # 3 chunks, accurate
38
39context = "\n\n---\n\n".join([doc.page_content for doc in reranked])
40print(f"\nContext ready: {len(context)} chars passing to LLM")
← PREV5. RetrievalNEXT →7. LLM Generation