From cb470e516b9711be37d344de87d5d377510f16c6 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=88=98=E8=88=AA=E5=AE=87?= <3364451258@qq.com> Date: Sat, 11 Jul 2026 19:42:20 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E6=B7=BB=E5=8A=A0=20HybridRetriever=20?= =?UTF-8?q?=E2=80=94=20BM25+=E5=90=91=E9=87=8F=E6=B7=B7=E5=90=88=E6=A3=80?= =?UTF-8?q?=E7=B4=A2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/core/retriever.py | 187 ++++++++++++++++++++++++++++++++++++++++ tests/test_retriever.py | 161 ++++++++++++++++++++++++++++++++++ 2 files changed, 348 insertions(+) create mode 100644 src/core/retriever.py create mode 100644 tests/test_retriever.py diff --git a/src/core/retriever.py b/src/core/retriever.py new file mode 100644 index 0000000..d69a3ab --- /dev/null +++ b/src/core/retriever.py @@ -0,0 +1,187 @@ +"""混合检索器 — BM25 关键词 + 向量语义联合检索.""" +from __future__ import annotations + +import logging +import re +from typing import TYPE_CHECKING, TypedDict + +from rank_bm25 import BM25Okapi + +from src.core.db import VectorDB +from src.core.embedder import Embedder + +if TYPE_CHECKING: + from src.core.reranker import Reranker + +logger = logging.getLogger("md-vector-db") + + +class SearchResult(TypedDict): + """检索结果类型.""" + + id: str + content: str + source_file: str + section_title: str + heading_level: int + chunk_index: int + score: float + bm25_score: float + vector_score: float + + +class HybridRetriever: + """BM25 + 向量混合检索器. + + Architecture: + 1. 向量检索取得 top_k * candidate_multiplier 候选 + 2. BM25 对候选打分 + 3. 加权融合排序(默认 0.7 向量 + 0.3 BM25) + 4. 返回 top_k 结果 + + 支持按 source_file 元数据过滤。 + """ + + def __init__( + self, + db: VectorDB, + embedder: Embedder, + collection_name: str, + bm25_weight: float = 0.3, + vector_candidate_multiplier: int = 3, + reranker: "Reranker | None" = None, + ): + self._db = db + self._embedder = embedder + self._collection_name = collection_name + self._bm25_weight = bm25_weight + self._vector_multiplier = vector_candidate_multiplier + self._reranker = reranker + self._bm25_index: BM25Okapi | None = None + self._bm25_docs: list[str] = [] + self._bm25_ids: list[str] = [] + + @property + def _collection(self): + return self._db.get_or_create_collection(self._collection_name) + + def search( + self, + query: str, + top_k: int = 10, + source_file: str | None = None, + ) -> list[dict]: + """混合检索. + + Args: + query: 查询文本 + top_k: 返回结果数量 + source_file: 可选,按源文件过滤 + + Returns: + 按混合分数降序排列的结果列表 + """ + # 1. 向量检索:多取候选 + vector_candidates = self._vector_search( + query, top_k * self._vector_multiplier, source_file + ) + if not vector_candidates: + return [] + + # 2. BM25 打分 + bm25_scored = self._bm25_rerank(query, vector_candidates) + + # 3. 分数融合 + fused = self._fuse_scores(bm25_scored, self._bm25_weight) + + # 4. 排序取 top_k + fused.sort(key=lambda x: x["score"], reverse=True) + + # 4.5 可选:Cross-Encoder 重排序 + if self._reranker is not None and len(fused) > 1: + fused = self._reranker.rerank(query, fused, top_k=top_k) + + return fused[:top_k] + + # -- 内部方法 -- + + def _vector_search( + self, query: str, n: int, source_file: str | None + ) -> list[dict]: + """向量检索取得候选.""" + embeddings = self._embedder.embed([query]) + if not embeddings: + return [] + query_embedding = embeddings[0] + where_filter = {"source_file": source_file} if source_file else None + results = self._collection.query( + query_embeddings=[query_embedding], + n_results=n, + where=where_filter, + include=["documents", "metadatas", "distances"], + ) + candidates = [] + if results["ids"] and results["ids"][0]: + for i, doc_id in enumerate(results["ids"][0]): + metadata = results["metadatas"][0][i] if results["metadatas"] else {} + distance = results["distances"][0][i] if results["distances"] else 0.0 + vector_score = max(0.0, round(1.0 - distance, 4)) + candidates.append({ + "id": doc_id, + "content": results["documents"][0][i] if results["documents"] else "", + "source_file": metadata.get("source_file", ""), + "section_title": metadata.get("section_title", ""), + "heading_level": metadata.get("heading_level", 0), + "chunk_index": metadata.get("chunk_index", 0), + "vector_score": vector_score, + }) + return candidates + + def _bm25_rerank(self, query: str, candidates: list[dict]) -> list[dict]: + """用 BM25 对候选列表重新打分.""" + if not candidates: + return candidates + tokenized_query = self._tokenize(query) + tokenized_candidates = [self._tokenize(c["content"]) for c in candidates] + bm25 = BM25Okapi(tokenized_candidates) + scores = bm25.get_scores(tokenized_query) + # 归一化 BM25 分数到 [0, 1] + max_score = max(scores) if max(scores) > 0 else 1.0 + for i, c in enumerate(candidates): + c["bm25_score"] = round(scores[i] / max_score, 4) + return candidates + + def _fuse_scores(self, items: list[dict], bm25_weight: float) -> list[dict]: + """加权融合向量分和 BM25 分.""" + vector_weight = 1.0 - bm25_weight + for item in items: + bm25_s = item.get("bm25_score", 0.0) + vec_s = item.get("vector_score", 0.0) + item["score"] = round(vec_s * vector_weight + bm25_s * bm25_weight, 4) + return items + + def _ensure_bm25_index(self) -> None: + """确保 BM25 索引已构建(从 collection 所有文档构建).""" + if self._bm25_index is not None: + return + all_data = self._collection.get(include=["documents", "metadatas"]) + if all_data and all_data["ids"]: + self._bm25_ids = all_data["ids"] + self._bm25_docs = all_data["documents"] or [] + tokenized = [self._tokenize(d) for d in self._bm25_docs] + self._bm25_index = BM25Okapi(tokenized) if tokenized else None + else: + self._bm25_ids = [] + self._bm25_docs = [] + self._bm25_index = None + + @staticmethod + def _tokenize(text: str) -> list[str]: + """简易中文+英文分词(按中文单字 + 英文单词拆分). + + 注意: 这是基础实现。生产环境建议集成 jieba 分词。 + """ + tokens = [] + for match in re.finditer(r"[a-zA-Z0-9]+|[一-鿿]|[^\s]", text): + tokens.append(match.group().lower()) + return tokens diff --git a/tests/test_retriever.py b/tests/test_retriever.py new file mode 100644 index 0000000..31e36fd --- /dev/null +++ b/tests/test_retriever.py @@ -0,0 +1,161 @@ +"""混合检索器测试.""" +import hashlib + +import pytest + +from src.core.retriever import HybridRetriever + + +class FakeEmbedder: + """模拟嵌入器 — 返回伪向量.""" + + @property + def dimension(self) -> int: + return 4 + + def embed(self, texts: list[str]) -> list[list[float]]: + result = [] + for t in texts: + h = hashlib.md5(t.encode()).digest() + vec = [float(b) / 255.0 for b in h[:4]] + result.append(vec) + return result + + +class FakeCollection: + """模拟 ChromaDB collection.""" + + def __init__(self): + self._docs: list[dict] = [] + + def add(self, ids, embeddings, documents, metadatas): + for i, doc_id in enumerate(ids): + self._docs.append({ + "id": doc_id, + "embedding": embeddings[i] if embeddings else [], + "document": documents[i], + "metadata": metadatas[i] if metadatas else {}, + }) + + def query(self, query_embeddings, n_results, where=None, include=None): + # 返回所有已入库文档(模拟向量检索) + n = min(n_results, len(self._docs)) + if n == 0: + return {"ids": [[]], "documents": [[]], "metadatas": [[]], "distances": [[]]} + ids_list = [d["id"] for d in self._docs[:n]] + docs_list = [d["document"] for d in self._docs[:n]] + metas_list = [d["metadata"] for d in self._docs[:n]] + dists_list = [0.2 + i * 0.05 for i in range(n)] # 伪距离 + return { + "ids": [ids_list], + "documents": [docs_list], + "metadatas": [metas_list], + "distances": [dists_list], + } + + def get(self, include=None): + return { + "ids": [d["id"] for d in self._docs], + "documents": [d["document"] for d in self._docs], + "metadatas": [d["metadata"] for d in self._docs], + } + + def count(self) -> int: + return len(self._docs) + + def delete(self, ids): + self._docs = [d for d in self._docs if d["id"] not in ids] + + +class FakeDB: + """模拟 VectorDB.""" + + def __init__(self): + self._collections: dict[str, FakeCollection] = {} + + def get_or_create_collection(self, name: str): + if name not in self._collections: + self._collections[name] = FakeCollection() + return self._collections[name] + + +@pytest.fixture +def retriever(): + db = FakeDB() + embedder = FakeEmbedder() + return HybridRetriever(db, embedder, "test", bm25_weight=0.3) + + +def test_search_returns_list(retriever): + """基本语义检索返回列表.""" + results = retriever.search("测试查询", top_k=5) + assert isinstance(results, list) + + +def test_search_with_source_filter(retriever): + """按 source_file 过滤.""" + results = retriever.search("查询", top_k=5, source_file="doc.md") + assert isinstance(results, list) + + +def test_search_top_k_bounds(retriever): + """top_k 在合理范围内.""" + for k in [1, 10, 50]: + results = retriever.search("test", top_k=k) + assert len(results) <= k + + +def test_bm25_index_built_from_collection(retriever): + """BM25 索引从 collection 文档构建.""" + coll = retriever._db.get_or_create_collection("test") + coll.add( + ids=["doc_0", "doc_1", "doc_2"], + embeddings=[[0.1] * 4, [0.2] * 4, [0.3] * 4], + documents=["Python 是一门编程语言", "Java 也是编程语言", "今天天气很好"], + metadatas=[ + {"source_file": "a.md"}, + {"source_file": "b.md"}, + {"source_file": "c.md"}, + ], + ) + retriever._bm25_index = None # 强制重建 + results = retriever.search("编程语言", top_k=2) + assert len(results) == 2 + assert any("编程语言" in r["content"] for r in results) + + +def test_hybrid_score_fusion(retriever): + """混合分数融合:向量分 + BM25 分加权.""" + coll = retriever._db.get_or_create_collection("test") + coll.add( + ids=["d0", "d1"], + embeddings=[[1.0] * 4, [0.5] * 4], + documents=["Docker 容器化部署指南", "Python 数据分析入门"], + metadatas=[{"source_file": "x.md"}, {"source_file": "y.md"}], + ) + retriever._bm25_index = None + results = retriever.search("Docker 部署", top_k=2) + assert len(results) >= 1 + assert "Docker" in results[0]["content"] + + +def test_empty_collection_returns_empty(retriever): + """空 collection 返回空列表.""" + results = retriever.search("查询", top_k=5) + assert results == [] + + +def test_metadata_in_results(retriever): + """结果中包含完整元数据.""" + coll = retriever._db.get_or_create_collection("test") + coll.add( + ids=["meta_test"], + embeddings=[[0.5] * 4], + documents=["带元数据的文档"], + metadatas=[{"source_file": "meta.md", "section_title": "第一章", "heading_level": 1}], + ) + retriever._bm25_index = None + results = retriever.search("元数据", top_k=1) + assert len(results) == 1 + assert results[0]["source_file"] == "meta.md" + assert results[0]["section_title"] == "第一章"