feat: 添加 HybridRetriever — BM25+向量混合检索

This commit is contained in:
2026-07-11 19:42:20 +08:00
parent 7b3c7d5323
commit cb470e516b
2 changed files with 348 additions and 0 deletions
+187
View File
@@ -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
+161
View File
@@ -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"] == "第一章"