Spaces:
Sleeping
Sleeping
File size: 2,411 Bytes
8eb009a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | from core.config import ProjectSettings
from serving.model_loader import _load_reranker_runtime
def test_load_reranker_runtime_none_backend() -> None:
settings = ProjectSettings(reranker_backend="none", use_cross_encoder=False)
runtime = _load_reranker_runtime(settings)
assert runtime.reranker is None
assert runtime.reranker_backend == "none"
assert runtime.fallback_used is False
def test_load_reranker_runtime_heuristic_backend() -> None:
settings = ProjectSettings(reranker_backend="heuristic", use_cross_encoder=False)
runtime = _load_reranker_runtime(settings)
assert runtime.reranker is not None
assert runtime.reranker_backend == "heuristic"
assert runtime.fallback_used is False
def test_load_reranker_runtime_cross_encoder_backend_uses_mock(monkeypatch) -> None:
class FakeCrossEncoderReranker:
def __init__(self, model_name: str) -> None:
self.model_name = model_name
self.backend_name = "cross-encoder"
def rank(self, claim, evidence_list): # noqa: ANN001
return list(reversed(evidence_list))
monkeypatch.setattr("serving.model_loader.CrossEncoderReranker", FakeCrossEncoderReranker)
settings = ProjectSettings(
reranker_backend="cross_encoder",
use_cross_encoder=True,
cross_encoder_model="fake-model",
)
runtime = _load_reranker_runtime(settings)
assert runtime.reranker is not None
assert getattr(runtime.reranker, "model_name", None) == "fake-model"
assert runtime.reranker_backend == "cross_encoder"
assert runtime.cross_encoder_model == "fake-model"
assert runtime.fallback_used is False
def test_load_reranker_runtime_cross_encoder_falls_back(monkeypatch) -> None:
class ExplodingCrossEncoderReranker:
def __init__(self, model_name: str) -> None: # noqa: ARG002
raise RuntimeError("download failed")
monkeypatch.setattr("serving.model_loader.CrossEncoderReranker", ExplodingCrossEncoderReranker)
settings = ProjectSettings(
reranker_backend="cross_encoder",
use_cross_encoder=True,
cross_encoder_model="fake-model",
)
runtime = _load_reranker_runtime(settings)
assert runtime.reranker is not None
assert runtime.reranker_backend == "heuristic"
assert runtime.cross_encoder_model == "fake-model"
assert runtime.fallback_used is True
|