feat: 添加 Cross-Encoder Reranker + 集成到 Searcher
This commit is contained in:
@@ -0,0 +1,69 @@
|
|||||||
|
"""Cross-Encoder 重排序器."""
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger("md-vector-db")
|
||||||
|
|
||||||
|
|
||||||
|
class Reranker:
|
||||||
|
"""使用 Cross-Encoder 模型对检索结果重排序.
|
||||||
|
|
||||||
|
默认模型: BAAI/bge-reranker-base(中文友好)
|
||||||
|
首次调用时懒加载模型。
|
||||||
|
"""
|
||||||
|
|
||||||
|
_DEFAULT_MODEL = "BAAI/bge-reranker-base"
|
||||||
|
|
||||||
|
def __init__(self, model_name: str | None = None):
|
||||||
|
self._model_name = model_name or self._DEFAULT_MODEL
|
||||||
|
self._model = None
|
||||||
|
|
||||||
|
def rerank(
|
||||||
|
self, query: str, candidates: list[dict], top_k: int = 10
|
||||||
|
) -> list[dict]:
|
||||||
|
"""对候选列表重排序.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
query: 原始查询
|
||||||
|
candidates: 候选结果列表(需含 "content" 字段)
|
||||||
|
top_k: 返回数量
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
按 cross-encoder 分数降序的结果列表
|
||||||
|
"""
|
||||||
|
if not candidates:
|
||||||
|
return []
|
||||||
|
|
||||||
|
self._ensure_model()
|
||||||
|
# 构建 (query, document) 对
|
||||||
|
pairs = [(query, c["content"]) for c in candidates]
|
||||||
|
|
||||||
|
try:
|
||||||
|
scores = self._model.predict(pairs, show_progress_bar=False)
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("Cross-Encoder 重排序失败: %s", e)
|
||||||
|
# 降级:保留原始顺序
|
||||||
|
return candidates[:top_k]
|
||||||
|
|
||||||
|
# 附加 rerank_score
|
||||||
|
for i, c in enumerate(candidates):
|
||||||
|
c["rerank_score"] = round(float(scores[i]), 4)
|
||||||
|
|
||||||
|
# 按 rerank_score 降序排列
|
||||||
|
candidates.sort(key=lambda x: x.get("rerank_score", 0), reverse=True)
|
||||||
|
|
||||||
|
# 返回 top_k,将 rerank_score 作为最终 score
|
||||||
|
result = candidates[:top_k]
|
||||||
|
for r in result:
|
||||||
|
r["score"] = r.get("rerank_score", r.get("score", 0))
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _ensure_model(self) -> None:
|
||||||
|
"""懒加载 Cross-Encoder 模型."""
|
||||||
|
if self._model is not None:
|
||||||
|
return
|
||||||
|
from sentence_transformers import CrossEncoder
|
||||||
|
|
||||||
|
logger.info("加载 Cross-Encoder 模型: %s", self._model_name)
|
||||||
|
self._model = CrossEncoder(self._model_name)
|
||||||
@@ -34,12 +34,17 @@ class Searcher:
|
|||||||
|
|
||||||
def _get_hybrid(self) -> HybridRetriever:
|
def _get_hybrid(self) -> HybridRetriever:
|
||||||
if self._hybrid is None:
|
if self._hybrid is None:
|
||||||
|
reranker = None
|
||||||
|
if self._search_config.enable_rerank:
|
||||||
|
from src.core.reranker import Reranker
|
||||||
|
reranker = Reranker()
|
||||||
self._hybrid = HybridRetriever(
|
self._hybrid = HybridRetriever(
|
||||||
self.db,
|
self.db,
|
||||||
self.embedder,
|
self.embedder,
|
||||||
self.collection_name,
|
self.collection_name,
|
||||||
bm25_weight=self._search_config.bm25_weight,
|
bm25_weight=self._search_config.bm25_weight,
|
||||||
vector_candidate_multiplier=self._search_config.candidate_multiplier,
|
vector_candidate_multiplier=self._search_config.candidate_multiplier,
|
||||||
|
reranker=reranker,
|
||||||
)
|
)
|
||||||
return self._hybrid
|
return self._hybrid
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,88 @@
|
|||||||
|
"""重排序器测试."""
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
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"]
|
||||||
Reference in New Issue
Block a user