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