Content hash: 0ec2eb7b1e43b741824ca32fdbec94b46a3fd9aa04a240ff88a6b983065aaa26
#!/usr/bin/env python3
"""Cross-encoder reranking demo: retrieve wide, rerank deep.
Shows the full two-stage pipeline with mock retrievers and RRF fusion,
plus optional real cross-encoder if sentence-transformers is installed.
"""
from __future__ import annotations
import sys
from dataclasses import dataclass
@dataclass
class Document:
id: str
text: str
score: float = 0.0
def bm25_retrieve(query: str, corpus: list[str], top_k: int = 20) -> list[Document]:
"""Mock BM25: simple term-overlap scoring."""
query_terms = set(query.lower().split())
scored = []
for i, doc in enumerate(corpus):
overlap = len(query_terms & set(doc.lower().split()))
if overlap > 0:
scored.append(Document(id=str(i), text=doc, score=overlap))
scored.sort(key=lambda d: d.score, reverse=True)
return scored[:top_k]
def dense_retrieve(query: str, corpus: list[str], top_k: int = 20) -> list[Document]:
"""Mock dense retriever: prefix-match heuristic for semantic similarity."""
scored = []
for i, doc in enumerate(corpus):
score = 0.0
for qw in query.lower().split():
for dw in doc.lower().split():
common = min(len(qw), len(dw))
if qw[:common] == dw[:common]:
score += common / max(len(qw), len(dw))
if score > 0:
scored.append(Document(id=str(i), text=doc, score=score))
scored.sort(key=lambda d: d.score, reverse=True)
return scored[:top_k]
def reciprocal_rank_fusion(ranked_lists: list[list[Document]], k: int = 60) -> list[Document]:
"""Fuse multiple ranked lists using RRF."""
scores: dict[str, tuple[float, Document]] = {}
for ranked in ranked_lists:
for rank, doc in enumerate(ranked):
rrf_score = 1.0 / (k + rank + 1)
if doc.id in scores:
prev_score, prev_doc = scores[doc.id]
scores[doc.id] = (prev_score + rrf_score, prev_doc)
else:
scores[doc.id] = (rrf_score, doc)
fused = sorted(scores.values(), key=lambda x: x[0], reverse=True)
return [doc for _, doc in fused]
def cross_encoder_rerank(query: str, candidates: list[Document]) -> list[tuple[Document, float]]:
"""Real or heuristic cross-encoder scoring of (query, doc) pairs."""
try:
from sentence_transformers import CrossEncoder
model = CrossEncoder("cross-encoder/ms-marco-MiniLM-L-6-v2")
pairs = [(query, doc.text) for doc in candidates]
scores = model.predict(pairs)
return [(doc, float(score)) for doc, score in zip(candidates, scores)]
except ImportError:
results = []
for doc in candidates:
score = 2.0 if query.lower() in doc.text.lower() else 0.0
score += len(set(query.lower().split()) & set(doc.text.lower().split()))
results.append((doc, score / max(len(query.split()), 1)))
return sorted(results, key=lambda x: x[1], reverse=True)
except Exception as e:
print(f"Warning: cross-encoder failed ({e})", file=sys.stderr)
return [(doc, 0.0) for doc in candidates]
def two_stage_retrieve(query: str, corpus: list[str], retrieval_k: int = 20, final_k: int = 3):
bm25_results = bm25_retrieve(query, corpus, top_k=retrieval_k)
dense_results = dense_retrieve(query, corpus, top_k=retrieval_k)
fused = reciprocal_rank_fusion([bm25_results, dense_results])
reranked = cross_encoder_rerank(query, fused[:retrieval_k])
return reranked[:final_k]
CORPUS = [
"Python is a popular programming language for data science and machine learning.",
"JavaScript runs in web browsers and is used for frontend development.",
"PostgreSQL is a powerful open-source relational database system.",
"Docker containers package applications with their dependencies.",
"Machine learning models require careful evaluation and testing.",
"Web development with Python often uses Django or FastAPI frameworks.",
"SQL databases like PostgreSQL support complex queries and transactions.",
"Frontend frameworks include React, Vue, and Angular for building UIs.",
"Data science pipelines often involve Python, pandas, and Jupyter notebooks.",
"Container orchestration with Kubernetes manages Docker deployments at scale.",
]
if __name__ == "__main__":
queries = [
"How to build web apps with Python?",
"database systems for SQL",
"machine learning evaluation",
]
for q in queries:
print(f"\nQuery: {q}")
results = two_stage_retrieve(q, CORPUS)
for rank, (doc, score) in enumerate(results, 1):
print(f" {rank}. [{score:.3f}] {doc.text[:80]}...")