Spaces:
Sleeping
Sleeping
feat: cross-encoder reranker with ONNX inference
Browse filesCo-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
- src/mediastorm/rag/reranker.py +47 -0
- tests/test_reranker.py +39 -0
src/mediastorm/rag/reranker.py
ADDED
|
@@ -0,0 +1,47 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
|
| 3 |
+
import numpy as np
|
| 4 |
+
import onnxruntime as ort
|
| 5 |
+
from tokenizers import Tokenizer
|
| 6 |
+
|
| 7 |
+
from mediastorm.config import RERANKER_MODEL_DIR
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class Reranker:
|
| 11 |
+
def __init__(self, model_dir: Path | str = RERANKER_MODEL_DIR):
|
| 12 |
+
model_dir = Path(model_dir)
|
| 13 |
+
self._tokenizer = Tokenizer.from_file(str(model_dir / "tokenizer.json"))
|
| 14 |
+
self._tokenizer.enable_padding()
|
| 15 |
+
self._tokenizer.enable_truncation(max_length=512)
|
| 16 |
+
self._session = ort.InferenceSession(
|
| 17 |
+
str(model_dir / "model.onnx"),
|
| 18 |
+
providers=["CPUExecutionProvider"],
|
| 19 |
+
)
|
| 20 |
+
|
| 21 |
+
def score_pairs(self, query: str, documents: list[str]) -> list[float]:
|
| 22 |
+
"""Score (query, document) pairs. Returns one relevance score per document."""
|
| 23 |
+
if not documents:
|
| 24 |
+
return []
|
| 25 |
+
|
| 26 |
+
encoded = self._tokenizer.encode_batch(
|
| 27 |
+
[(query, doc) for doc in documents],
|
| 28 |
+
)
|
| 29 |
+
|
| 30 |
+
input_ids = np.array([e.ids for e in encoded], dtype=np.int64)
|
| 31 |
+
attention_mask = np.array([e.attention_mask for e in encoded], dtype=np.int64)
|
| 32 |
+
|
| 33 |
+
feeds = {"input_ids": input_ids, "attention_mask": attention_mask}
|
| 34 |
+
|
| 35 |
+
# Add token_type_ids if the model expects them
|
| 36 |
+
input_names = [inp.name for inp in self._session.get_inputs()]
|
| 37 |
+
if "token_type_ids" in input_names:
|
| 38 |
+
token_type_ids = np.array([e.type_ids for e in encoded], dtype=np.int64)
|
| 39 |
+
feeds["token_type_ids"] = token_type_ids
|
| 40 |
+
|
| 41 |
+
outputs = self._session.run(None, feeds)
|
| 42 |
+
logits = outputs[0]
|
| 43 |
+
|
| 44 |
+
# Cross-encoder output: (batch_size, 1) logit — squeeze to flat list
|
| 45 |
+
if logits.ndim == 2 and logits.shape[1] == 1:
|
| 46 |
+
return logits[:, 0].tolist()
|
| 47 |
+
return logits.flatten().tolist()
|
tests/test_reranker.py
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import pytest
|
| 2 |
+
from mediastorm.rag.reranker import Reranker
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
@pytest.fixture(scope="module")
|
| 6 |
+
def reranker():
|
| 7 |
+
return Reranker()
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def test_score_pairs_returns_one_score_per_document(reranker):
|
| 11 |
+
scores = reranker.score_pairs("climate change", ["global warming effects", "jazz music history"])
|
| 12 |
+
assert len(scores) == 2
|
| 13 |
+
assert all(isinstance(s, float) for s in scores)
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
def test_relevant_document_scores_higher(reranker):
|
| 17 |
+
scores = reranker.score_pairs(
|
| 18 |
+
"climate change",
|
| 19 |
+
["global warming and rising sea levels", "jazz musicians in New York"],
|
| 20 |
+
)
|
| 21 |
+
assert scores[0] > scores[1], f"Relevant doc should score higher: {scores}"
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def test_score_pairs_deterministic(reranker):
|
| 25 |
+
s1 = reranker.score_pairs("test query", ["test document"])
|
| 26 |
+
s2 = reranker.score_pairs("test query", ["test document"])
|
| 27 |
+
assert s1[0] == pytest.approx(s2[0], abs=1e-5)
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def test_score_pairs_handles_long_text(reranker):
|
| 31 |
+
long_text = "word " * 1000
|
| 32 |
+
scores = reranker.score_pairs("query", [long_text])
|
| 33 |
+
assert len(scores) == 1
|
| 34 |
+
assert isinstance(scores[0], float)
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
def test_score_pairs_empty_list(reranker):
|
| 38 |
+
scores = reranker.score_pairs("query", [])
|
| 39 |
+
assert scores == []
|