diff --git a/src/core/reranker.py b/src/core/reranker.py new file mode 100644 index 0000000..4d19c38 --- /dev/null +++ b/src/core/reranker.py @@ -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) diff --git a/src/core/search.py b/src/core/search.py index 17f4cda..adb6040 100644 --- a/src/core/search.py +++ b/src/core/search.py @@ -34,12 +34,17 @@ class Searcher: def _get_hybrid(self) -> HybridRetriever: 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.db, self.embedder, self.collection_name, bm25_weight=self._search_config.bm25_weight, vector_candidate_multiplier=self._search_config.candidate_multiplier, + reranker=reranker, ) return self._hybrid diff --git a/tests/test_reranker.py b/tests/test_reranker.py new file mode 100644 index 0000000..ae0550e --- /dev/null +++ b/tests/test_reranker.py @@ -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"]