Files
QMDSearch/tests/test_ranker.py
T
kplam 51dc8dc4f6 Initial commit: QMDSearch 分层信息检索服务
- FastAPI + Qdrant + Redis + Ollama 技术栈
- L1→L2→L3→chunk 四层分层检索(dense + sparse RRF 融合)
- 文档三级总结与 2.5 级回退
- query 解析路由与分类
- /admin 管理页面
2026-07-29 21:24:40 +08:00

80 lines
3.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""RRF 融合与截断的单元测试"""
import pytest
from qdrant_client import models
from app.core.ranker import finalize, rrf_fuse
def _point(pid: str, score: float = 1.0) -> models.ScoredPoint:
return models.ScoredPoint(id=pid, version=0, score=score, payload={}, vector=None)
class TestRrfFuse:
"""rrf_fuse 的排序 / 去重 / 空输入 / k 参数行为"""
def test_empty_input(self):
"""空列表输入返回 []"""
assert rrf_fuse([]) == []
assert rrf_fuse([[], []]) == []
def test_single_list_fused_scores(self):
"""单路结果按 rank 计算 1/(k+rank) 融合分并降序返回"""
fused = rrf_fuse([[_point("a"), _point("b"), _point("c")]], k=60)
assert [p.id for p in fused] == ["a", "b", "c"]
assert fused[0].score == pytest.approx(1 / 61)
assert fused[1].score == pytest.approx(1 / 62)
assert fused[2].score == pytest.approx(1 / 63)
def test_dedup_merges_scores(self):
"""同一 point id 出现在多路结果中,只保留一条且分数累加"""
list1 = [_point("a"), _point("b")]
list2 = [_point("b"), _point("c")]
fused = rrf_fuse([list1, list2], k=60)
# b 在两路中分别 rank2/rank1,融合分最高;a 与 c 各 1/61
assert [p.id for p in fused] == ["b", "a", "c"]
assert fused[0].score == pytest.approx(1 / 62 + 1 / 61)
assert fused[1].score == pytest.approx(1 / 61)
assert fused[2].score == pytest.approx(1 / 62)
def test_k_affects_ranking(self):
"""k 参数改变融合权重,可改变最终排序
A 排名 [1, 10]B 排名 [5, 5]
- k=60 时 B 总分更高(两路均衡占优)
- k=1 时 A 总分更高(头部 rank 权重被放大)
"""
l1 = [_point("a"), *[_point(f"n{i}") for i in range(3)], _point("b"), *[_point(f"m{i}") for i in range(5)]]
l2 = [*[_point(f"x{i}") for i in range(4)], _point("b"), *[_point(f"y{i}") for i in range(4)], _point("a")]
# l1a rank1b rank5l2b rank5a rank10
fused_k60 = rrf_fuse([l1, l2], k=60)
assert fused_k60[0].id == "b"
fused_k1 = rrf_fuse([l1, l2], k=1)
assert fused_k1[0].id == "a"
def test_result_is_new_objects(self):
"""返回的是写回融合分的新对象,不修改原 point"""
original = _point("a", score=0.99)
fused = rrf_fuse([[original]], k=60)
assert fused[0].score != 0.99
assert original.score == 0.99
class TestFinalize:
"""finalize 截断行为"""
def test_truncate(self):
points = [_point(str(i)) for i in range(10)]
result = finalize(points, 5)
assert [p.id for p in result] == ["0", "1", "2", "3", "4"]
def test_truncate_larger_than_input(self):
"""final_k 大于列表长度时返回全部"""
points = [_point("a"), _point("b")]
assert finalize(points, 5) == points
def test_empty(self):
assert finalize([], 5) == []