feat: 添加 HybridRetriever — BM25+向量混合检索
This commit is contained in:
@@ -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
|
||||
@@ -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"] == "第一章"
|
||||
Reference in New Issue
Block a user