Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 12 additions & 1 deletion app.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from src.retrieval.bm25_index import BM25Index
from src.retrieval.embeddings import make_ollama_embed_fn
from src.retrieval.hybrid import HybridRetriever
from src.retrieval.reranker import Reranker
from src.retrieval.vector_store import VectorStore

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -83,6 +84,10 @@ async def on_chat_start():
)
cl.user_session.set("retriever", retriever)

# Load cross-encoder reranker
reranker = Reranker(model_name=config.retrieval.reranker_model)
cl.user_session.set("reranker", reranker)

if doc_count > 0:
await cl.Message(
content=f"{doc_count} chunks indexed. Ask me anything!"
Expand All @@ -109,7 +114,13 @@ async def on_message(message: cl.Message):
).send()
return

results = retriever.retrieve(message.content)
candidates = retriever.retrieve(message.content)

# Rerank candidates with cross-encoder
reranker = cl.user_session.get("reranker")
results = reranker.rerank(
message.content, candidates, top_k=config.retrieval.rerank_top_k
)

# Stream the answer token by token
msg = cl.Message(content="")
Expand Down
40 changes: 40 additions & 0 deletions src/retrieval/reranker.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
"""Cross-encoder reranker for search result refinement.

Not via LangChain: uses sentence-transformers CrossEncoder directly
for full control over scoring and result construction.
"""

from sentence_transformers import CrossEncoder

from src.models import SearchResult


class Reranker:
"""Scores query-chunk pairs with a cross-encoder and returns top results."""

def __init__(self, model_name: str) -> None:
self._model = CrossEncoder(model_name)

def rerank(
self,
query: str,
results: list[SearchResult],
top_k: int,
) -> list[SearchResult]:
"""Score all results against the query and return top_k by relevance."""
if not results:
return []

pairs = [(query, r.text) for r in results]
scores = self._model.predict(pairs)

scored = sorted(
zip(scores, results),
key=lambda x: x[0],
reverse=True,
)

return [
SearchResult(text=r.text, metadata=r.metadata, distance=float(score))
for score, r in scored[:top_k]
]
132 changes: 132 additions & 0 deletions tests/test_reranker.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,132 @@
"""Tests for cross-encoder reranker."""

from unittest.mock import MagicMock, patch

import pytest

from src.models import SearchResult
from src.retrieval.reranker import Reranker


@pytest.fixture
def sample_results() -> list[SearchResult]:
"""Create search results where order should change after reranking."""
return [
SearchResult(
text="The capital of France is Paris.",
metadata={"filename": "geography.pdf", "page_number": 1},
distance=0.3,
),
SearchResult(
text="Python is a programming language.",
metadata={"filename": "tech.pdf", "page_number": 5},
distance=0.4,
),
SearchResult(
text="Paris has many famous landmarks including the Eiffel Tower.",
metadata={"filename": "travel.pdf", "page_number": 10},
distance=0.5,
),
]


class TestReranker:
"""Tests for the Reranker class."""

@patch("src.retrieval.reranker.CrossEncoder")
def test_rerank_returns_top_k_results(
self, mock_cross_encoder_cls, sample_results
):
"""Reranker returns at most top_k results."""
mock_model = MagicMock()
mock_model.predict.return_value = [0.9, 0.1, 0.7]
mock_cross_encoder_cls.return_value = mock_model

reranker = Reranker(model_name="test-model")
reranked = reranker.rerank("What is Paris?", sample_results, top_k=2)

assert len(reranked) == 2

@patch("src.retrieval.reranker.CrossEncoder")
def test_rerank_reorders_by_score(
self, mock_cross_encoder_cls, sample_results
):
"""Reranker reorders results by cross-encoder score (highest first)."""
mock_model = MagicMock()
# Scores: first=0.9, second=0.1, third=0.7
mock_model.predict.return_value = [0.9, 0.1, 0.7]
mock_cross_encoder_cls.return_value = mock_model

reranker = Reranker(model_name="test-model")
reranked = reranker.rerank("What is Paris?", sample_results, top_k=3)

assert reranked[0].text == "The capital of France is Paris."
assert reranked[1].text == "Paris has many famous landmarks including the Eiffel Tower."
assert reranked[2].text == "Python is a programming language."

@patch("src.retrieval.reranker.CrossEncoder")
def test_rerank_updates_distance_to_score(
self, mock_cross_encoder_cls, sample_results
):
"""Reranked results have cross-encoder scores as distance."""
mock_model = MagicMock()
mock_model.predict.return_value = [0.9, 0.1, 0.7]
mock_cross_encoder_cls.return_value = mock_model

reranker = Reranker(model_name="test-model")
reranked = reranker.rerank("What is Paris?", sample_results, top_k=3)

assert reranked[0].distance == 0.9
assert reranked[1].distance == 0.7
assert reranked[2].distance == 0.1

@patch("src.retrieval.reranker.CrossEncoder")
def test_rerank_passes_query_chunk_pairs(
self, mock_cross_encoder_cls, sample_results
):
"""Reranker passes correct (query, text) pairs to the model."""
mock_model = MagicMock()
mock_model.predict.return_value = [0.5, 0.5, 0.5]
mock_cross_encoder_cls.return_value = mock_model

reranker = Reranker(model_name="test-model")
reranker.rerank("What is Paris?", sample_results, top_k=3)

pairs = mock_model.predict.call_args[0][0]
assert len(pairs) == 3
assert pairs[0] == ("What is Paris?", "The capital of France is Paris.")

@patch("src.retrieval.reranker.CrossEncoder")
def test_rerank_empty_input(self, mock_cross_encoder_cls):
"""Reranker handles empty input gracefully."""
mock_model = MagicMock()
mock_cross_encoder_cls.return_value = mock_model

reranker = Reranker(model_name="test-model")
reranked = reranker.rerank("query", [], top_k=5)

assert reranked == []
mock_model.predict.assert_not_called()

@patch("src.retrieval.reranker.CrossEncoder")
def test_rerank_top_k_larger_than_input(
self, mock_cross_encoder_cls, sample_results
):
"""When top_k exceeds input size, return all results."""
mock_model = MagicMock()
mock_model.predict.return_value = [0.9, 0.1, 0.7]
mock_cross_encoder_cls.return_value = mock_model

reranker = Reranker(model_name="test-model")
reranked = reranker.rerank("query", sample_results, top_k=10)

assert len(reranked) == 3

@patch("src.retrieval.reranker.CrossEncoder")
def test_model_loaded_with_correct_name(self, mock_cross_encoder_cls):
"""CrossEncoder is initialised with the configured model name."""
mock_cross_encoder_cls.return_value = MagicMock()

Reranker(model_name="bge-reranker-v2-m3")

mock_cross_encoder_cls.assert_called_once_with("bge-reranker-v2-m3")
Loading