haiku.rag/haiku_rag_slim/haiku/rag/reranking/cross_encoder.py
Yiorgis Gozadinos 72ef18e39d
Reject unknown and out-of-range configuration values
Every section inherited plain BaseModel, so unknown keys were dropped
silently: providers.docling_serve.timeout was documented for months while
being ignored, and a typo in any setting took the default. Sections now
derive from ConfigModel, which forbids extras, so a stale or misspelled key
fails with its path. This already found search.context_radius in a live app
config and providers.vllm in soliplex's example.

converter, chunker and chunker_type are Literals. Sizes, limits,
dimensions, token budgets, attempt counts and breaker thresholds must be
positive; retention, delays, intervals and cooldowns non-negative;
similarity_threshold within 0-1; port within 0-65535. port 0 keeps its
OS-assigned meaning and worker_count allows 0 for an API-and-reaper-only
process.

get_reranker caught ImportError and returned None, so a configured reranker
whose extra was missing silently disappeared. It now propagates.
raise_missing_extra names the install command and re-raises when the failure
came from inside an installed package, so a broken transitive import is not
reported as a missing one. zeroentropy imported bare and now guards like the
others.

The haiku.rag package declares the jina extra. jina-local already worked
there through cross-encoder's transitive transformers and torch; the
resolved package set is unchanged, but the support is now promised rather
than inherited.

Provider fields stay unconstrained: get_model ends in a pass-through to
pydantic-ai for any provider it supports, so a Literal there would reject
valid configurations.
2026-08-19 15:32:51 +03:00

45 lines
1.5 KiB
Python

import asyncio
import math
from haiku.rag.utils import raise_missing_extra
try:
import torch
from sentence_transformers import CrossEncoder
except ModuleNotFoundError as e: # pragma: no cover
if e.name not in ("torch", "sentence_transformers"):
raise
raise_missing_extra(e.name, "cross-encoder", 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
]