88 lines
2.7 KiB
Python
88 lines
2.7 KiB
Python
"""重排序器测试."""
|
|
|
|
from src.core.reranker import Reranker
|
|
|
|
|
|
class FakeCrossEncoder:
|
|
"""模拟 Cross-Encoder 模型."""
|
|
|
|
def predict(self, pairs, **kwargs):
|
|
# 包含"重要"的 pair 分数高
|
|
scores = []
|
|
for pair in pairs:
|
|
score = 5.0 if "重要" in pair[1] else 1.0
|
|
scores.append(score)
|
|
return scores
|
|
|
|
|
|
def test_reranker_returns_same_count():
|
|
"""重排序不改变结果数量."""
|
|
reranker = Reranker(model_name="test-model")
|
|
reranker._model = FakeCrossEncoder()
|
|
candidates = [
|
|
{"content": "普通文档", "score": 0.8},
|
|
{"content": "重要文档", "score": 0.6},
|
|
{"content": "另一个普通", "score": 0.7},
|
|
]
|
|
result = reranker.rerank("查询", candidates, top_k=3)
|
|
assert len(result) == 3
|
|
|
|
|
|
def test_reranker_promotes_relevant():
|
|
"""重排序将更相关的内容提前."""
|
|
reranker = Reranker(model_name="test-model")
|
|
reranker._model = FakeCrossEncoder()
|
|
candidates = [
|
|
{"content": "普通 A", "score": 0.9},
|
|
{"content": "重要内容在这里", "score": 0.5},
|
|
{"content": "普通 B", "score": 0.7},
|
|
]
|
|
result = reranker.rerank("查询", candidates, top_k=3)
|
|
assert "重要" in result[0]["content"]
|
|
|
|
|
|
def test_reranker_truncates_to_top_k():
|
|
"""rerank 截断到指定的 top_k."""
|
|
reranker = Reranker(model_name="test-model")
|
|
reranker._model = FakeCrossEncoder()
|
|
candidates = [
|
|
{"content": f"文档{i}", "score": 0.9 - i * 0.1}
|
|
for i in range(20)
|
|
]
|
|
result = reranker.rerank("查询", candidates, top_k=5)
|
|
assert len(result) == 5
|
|
|
|
|
|
def test_reranker_empty_input():
|
|
"""空输入返回空列表."""
|
|
reranker = Reranker(model_name="test-model")
|
|
result = reranker.rerank("查询", [], top_k=5)
|
|
assert result == []
|
|
|
|
|
|
def test_reranker_preserves_metadata():
|
|
"""重排序保留文档元数据."""
|
|
reranker = Reranker(model_name="test-model")
|
|
reranker._model = FakeCrossEncoder()
|
|
candidates = [
|
|
{
|
|
"content": "带元数据的文档",
|
|
"score": 0.5,
|
|
"source_file": "meta.md",
|
|
"section_title": "第一章",
|
|
}
|
|
]
|
|
result = reranker.rerank("查询", candidates, top_k=1)
|
|
assert result[0]["source_file"] == "meta.md"
|
|
assert result[0]["section_title"] == "第一章"
|
|
|
|
|
|
def test_reranker_score_replaced_with_rerank():
|
|
"""重排序后 score 更新为 rerank_score."""
|
|
reranker = Reranker(model_name="test-model")
|
|
reranker._model = FakeCrossEncoder()
|
|
candidates = [{"content": "测试", "score": 0.5}]
|
|
result = reranker.rerank("查询", candidates, top_k=1)
|
|
assert "rerank_score" in result[0]
|
|
assert result[0]["score"] == result[0]["rerank_score"]
|