mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-12 23:12:23 -06:00
f6bae70ea6
Phase 2 of BM25 reranking (follows #627). Makes the rerank_bm25_threshold floor usable across reranker models and adds tooling to pick it. - normalize_scores (rerank.py): map a rerank batch into a 0-1 relevance probability -- sigmoid when any score falls outside [0,1] (logit endpoints like bge/TEI), identity otherwise (Cohere/Jina/Qwen already 0-1). Applied in the _bm25_reranker closure AND calibration so the threshold means the same on every endpoint. Monotonic, so ranking order is unchanged. - rerank_calibrate.py + `turnstone-admin rerank-calibrate [--apply]`: probe the endpoint with labelled relevant/irrelevant groups, normalise, and recommend a recall-biased floor -- or report "no clean separation" (a mis-served/weak reranker). A warmup loop absorbs a cold endpoint's first-request compile so calibration doesn't time out. Validated live against Qwen3-Reranker 0.6B and 4B: the calibrated floor differs sharply per model (~0.95 vs ~0.33 for the same task) -- exactly why per-endpoint calibration exists. - rerank_config.py: extract resolve_rerank_client_from(config_store, registry); the alias/url precedence now lives in one place, shared by ChatSession (which delegates) and the CLI. - tools.rerank_instruction (config + setting + client): wrap the query as <Instruct>:/<Query>: for instruction-aware rerankers (Qwen3) on endpoints that don't apply the model's own chat template. Docs note the critical vLLM serving detail: Qwen3-Reranker needs --chat-template or its scores are near-random and reranking hurts retrieval. Negative-tested: normalize sigmoid/identity branches, closure-normalises-before- floor, calibration separation/recall-bias/warmup-absorbs-cold-start, the CLI apply/no-apply/no-separation paths, and instruction query-wrapping through the real httpx boundary.
96 lines
4.0 KiB
Python
96 lines
4.0 KiB
Python
"""Calibration core — probe a fake reranker and recommend a 0-1 floor."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from turnstone.core.rerank import RerankHit
|
|
from turnstone.core.rerank_calibrate import _GAP_FRACTION, _PROBE_SET, _build_result, calibrate
|
|
|
|
|
|
class _ScriptedClient:
|
|
"""RerankClient stub: scores doc 0 (the relevant one) at ``r``, the rest ``i``.
|
|
|
|
Drives the real ``calibrate`` loop through the ``RerankClient`` seam — the
|
|
relevant doc is always position 0 of the documents calibrate sends.
|
|
"""
|
|
|
|
def __init__(self, r: float, i: float) -> None:
|
|
self._r, self._i = r, i
|
|
|
|
def rerank(
|
|
self, query: str, documents: list[str], *, top_n: int | None = None
|
|
) -> list[RerankHit]:
|
|
assert top_n is None # calibration must request every doc's score
|
|
return [RerankHit(index=0, score=self._r)] + [
|
|
RerankHit(index=idx, score=self._i) for idx in range(1, len(documents))
|
|
]
|
|
|
|
|
|
class _FlakyClient:
|
|
"""Fails the first ``cold`` calls (cold-start compile), then scores normally."""
|
|
|
|
def __init__(self, cold: int, r: float, i: float) -> None:
|
|
self.calls = 0
|
|
self.cold = cold
|
|
self._r, self._i = r, i
|
|
|
|
def rerank(
|
|
self, query: str, documents: list[str], *, top_n: int | None = None
|
|
) -> list[RerankHit]:
|
|
self.calls += 1
|
|
if self.calls <= self.cold:
|
|
raise RuntimeError("cold endpoint (compiling)")
|
|
return [RerankHit(index=0, score=self._r)] + [
|
|
RerankHit(index=idx, score=self._i) for idx in range(1, len(documents))
|
|
]
|
|
|
|
|
|
class TestCalibrate:
|
|
def test_warmup_absorbs_cold_start(self):
|
|
# First 2 calls fail (compile); warmup consumes them so the probe loop is
|
|
# warm and calibration still succeeds.
|
|
c = _FlakyClient(cold=2, r=0.9, i=0.1)
|
|
res = calibrate(c, model="m")
|
|
assert res.separated
|
|
assert c.calls > 2 # warmup absorbed the cold calls before the probes ran
|
|
|
|
def test_probability_scale_clean_separation(self):
|
|
res = calibrate(_ScriptedClient(0.9, 0.1), model="m")
|
|
assert res.raw_scale == "probability (0-1)" # already 0-1 -> identity
|
|
assert res.separated
|
|
# gap (0.1, 0.9); _GAP_FRACTION in from the irrelevant edge.
|
|
assert res.suggested_threshold == round(0.1 + _GAP_FRACTION * 0.8, 4)
|
|
assert res.irrelevant_max < res.suggested_threshold < res.relevant_min
|
|
assert res.n_relevant == len(_PROBE_SET)
|
|
assert res.n_irrelevant == len(_PROBE_SET) * (len(_PROBE_SET) - 1)
|
|
|
|
def test_logit_scale_normalized_then_separated(self):
|
|
# Out-of-[0,1] raw scores -> sigmoid -> a 0-1 floor regardless of scale.
|
|
res = calibrate(_ScriptedClient(5.0, -2.0), model="m")
|
|
assert "logit" in res.raw_scale
|
|
assert res.separated
|
|
assert res.suggested_threshold is not None
|
|
assert 0.0 < res.suggested_threshold < 1.0
|
|
assert res.irrelevant_max < res.suggested_threshold < res.relevant_min
|
|
# all reported score fields live in the normalised 0-1 space
|
|
assert 0.0 <= res.irrelevant_min <= res.relevant_max <= 1.0
|
|
|
|
def test_overlap_reports_no_separation(self):
|
|
# relevant 0.4 <= irrelevant 0.6 -> not separable, no recommendation.
|
|
res = calibrate(_ScriptedClient(0.4, 0.6), model="m")
|
|
assert not res.separated
|
|
assert res.suggested_threshold is None
|
|
|
|
def test_recall_bias_floor_below_lowest_relevant(self):
|
|
# The floor must never exceed the lowest relevant score (no false drops).
|
|
res = calibrate(_ScriptedClient(0.55, 0.45), model="m")
|
|
assert res.separated
|
|
assert res.suggested_threshold is not None
|
|
assert res.suggested_threshold < res.relevant_min
|
|
|
|
def test_empty_scores_is_no_separation(self):
|
|
# A broken endpoint that scores nothing -> health-check fail, no floor.
|
|
res = _build_result("m", "unknown (no scores)", [], [])
|
|
assert not res.separated
|
|
assert res.suggested_threshold is None
|
|
assert res.raw_scale == "unknown (no scores)"
|