haiku.rag/haiku_rag_slim/haiku/rag/reranking/cross_encoder.py
Yiorgis Gozadinos d63199d96d
Keep cross-encoder rerank scores apart when they saturate
mxbai-rerank-base-v2 ships a Sigmoid activation and evaluates it in bf16, so
every strongly-relevant candidate rounds to exactly 1.0. Ties then leave the
order to the stable sort, which preserves the incoming hybrid ranking: on 100
t2_finqa retrieval cases the reranker scored MAP 0.661 against 0.659 with no
reranker at all, and 0.742 once the scores separate.

Ask the model for logits and apply the sigmoid here, where it runs in float64.
Scores stay 0-1, matching the cohere, vllm and zeroentropy rerankers.

Also drop the remaining pyright references; the project type-checks with ty.
2026-08-07 14:07:53 +03:00

44 lines
1.5 KiB
Python

import asyncio
import math
try:
import torch
from sentence_transformers import CrossEncoder
except ImportError as e: # pragma: no cover
raise ImportError(
"sentence-transformers is not installed. Install it with "
"`pip install sentence-transformers` or use the cross-encoder optional dependency."
) from e
from haiku.rag.reranking.base import RerankerBase
from haiku.rag.store.models.chunk import Chunk
class CrossEncoderReranker(RerankerBase):
"""Reranker for any sentence-transformers CrossEncoder model.
Loads the model in-process. Pass any HuggingFace cross-encoder reranker
as ``model`` (e.g. ``BAAI/bge-reranker-v2-m3``, ``Qwen/Qwen3-Reranker-0.6B``,
``cross-encoder/ms-marco-MiniLM-L-6-v2``).
"""
def __init__(self, model: str):
self._model = model
self._reranker = CrossEncoder(model)
async def _rerank(
self, query: str, chunks: list[Chunk], top_n: int = 10
) -> list[tuple[Chunk, float]]:
documents = [chunk.content for chunk in chunks]
# Ask for logits and squash them here: the model's own sigmoid runs in
# bf16, where saturated scores round onto identical values and leave the
# order of the top candidates to the sort.
rankings = await asyncio.to_thread(
lambda: self._reranker.rank(
query, documents, top_k=top_n, activation_fn=torch.nn.Identity()
)
)
return [
(chunks[r["corpus_id"]], 1.0 / (1.0 + math.exp(-r["score"])))
for r in rankings
]