算法核心特点
双重检索:同时支持语义检索(向量)和关键词检索(BM25)
智能融合:支持多种融合策略,包括固定权重和熵权法
得分归一化:统一不同检索器的得分尺度
重排序:使用交叉编码器或规则进行精排
动态权重:根据检索结果质量动态调整权重
中文优化:支持jieba分词和中文停用词过滤
可扩展性:支持自定义分词器和融合策略
import numpy as np
from typing import List, Dict, Any, Optional, Tuple, Set
from collections import defaultdict, Counter
import math
import jieba
import jieba.posseg as pseg
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.metrics.pairwise import cosine_similarity
import json
import logging
from dataclasses import dataclass, field
from abc import ABC, abstractmethod
import re
# 配置日志
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class ChineseTokenizer:
"""中文分词器"""
def __init__(self, use_jieba: bool = True):
self.use_jieba = use_jieba
if use_jieba:
# 加载自定义词典(可选)
# jieba.load_userdict("custom_dict.txt")
pass
def tokenize(self, text: str) -> List[str]:
"""分词"""
if self.use_jieba:
# 使用jieba分词
words = jieba.lcut(text)
# 过滤停用词和单字符
stopwords = set(['的', '了', '在', '是', '我', '有', '和', '就', '不', '人', '都', '一', '一个', '上', '也', '很', '到', '说', '要', '去', '你', '会', '着', '没有', '看', '好', '自己', '这'])
return [w for w in words if w not in stopwords and len(w) > 1]
else:
# 简单的按字符分割
return list(text)
@dataclass
class Document:
"""文档类"""
doc_id: str
content: str
tokens: List[str] = field(default_factory=list)
embedding: Optional[np.ndarray] = None
metadata: Dict[str, Any] = field(default_factory=dict)
class BM25Retriever:
"""
BM25关键词检索器
实现经典的BM25算法,用于关键词搜索
"""
def __init__(
self,
tokenizer: Optional[ChineseTokenizer] = None,
k1: float = 1.5,
b: float = 0.75,
epsilon: float = 0.25
):
"""
初始化BM25检索器
Args:
tokenizer: 分词器
k1: BM25参数,控制词频饱和度
b: BM25参数,控制文档长度归一化
epsilon: 平滑参数
"""
self.tokenizer = tokenizer or ChineseTokenizer()
self.k1 = k1
self.b = b
self.epsilon = epsilon
self.documents: List[Document] = []
self.doc_freqs: Dict[str, int] = defaultdict(int) # 文档频率
self.idf: Dict[str, float] = {}
self.doc_len: List[int] = []
self.avg_doc_len: float = 0.0
self.total_docs: int = 0
self.is_built = False
def add_documents(self, documents: List[Dict[str, Any]]):
"""
添加文档到索引
Args:
documents: 文档列表,每个文档包含 'doc_id', 'content', 'metadata'
"""
for doc in documents:
doc_id = doc.get('doc_id', f"doc_{len(self.documents)}")
content = doc.get('content', '')
# 分词
tokens = self.tokenizer.tokenize(content)
# 创建文档对象
document = Document(
doc_id=doc_id,
content=content,
tokens=tokens,
metadata=doc.get('metadata', {})
)
self.documents.append(document)
# 重建索引
self._build_index()
logger.info(f"BM25索引已构建,包含 {len(self.documents)} 个文档")
def _build_index(self):
"""构建BM25索引"""
if not self.documents:
return
self.total_docs = len(self.documents)
self.doc_len = [len(doc.tokens) for doc in self.documents]
self.avg_doc_len = sum(self.doc_len) / self.total_docs if self.total_docs > 0 else 0
# 计算文档频率
self.doc_freqs.clear()
for doc in self.documents:
unique_tokens = set(doc.tokens)
for token in unique_tokens:
self.doc_freqs[token] += 1
# 计算IDF
self.idf.clear()
for token, freq in self.doc_freqs.items():
# 使用BM25的IDF公式
idf = math.log((self.total_docs - freq + 0.5) / (freq + 0.5) + 1)
self.idf[token] = idf
self.is_built = True
logger.info(f"BM25索引构建完成,词汇量: {len(self.idf)}")
def search(self, query: str, top_k: int = 10) -> List[Dict[str, Any]]:
"""
执行BM25搜索
Args:
query: 查询文本
top_k: 返回结果数量
Returns:
搜索结果列表
"""
if not self.is_built:
logger.warning("BM25索引未构建")
return []
# 查询分词
query_tokens = self.tokenizer.tokenize(query)
if not query_tokens:
return []
# 计算每个文档的BM25得分
scores = []
for idx, doc in enumerate(self.documents):
score = self._compute_bm25_score(query_tokens, idx)
scores.append({
'doc_id': doc.doc_id,
'content': doc.content,
'score': score,
'index': idx,
'metadata': doc.metadata
})
# 按得分排序
scores.sort(key=lambda x: x['score'], reverse=True)
# 返回top-k
return scores[:top_k]
def _compute_bm25_score(self, query_tokens: List[str], doc_idx: int) -> float:
"""
计算单个文档的BM25得分
Args:
query_tokens: 查询分词
doc_idx: 文档索引
Returns:
BM25得分
"""
score = 0.0
doc_len = self.doc_len[doc_idx]
# 统计文档中词频
doc_token_counts = Counter(self.documents[doc_idx].tokens)
for token in query_tokens:
if token not in self.idf:
continue
# 词频
tf = doc_token_counts.get(token, 0)
# 计算BM25得分
idf = self.idf[token]
# BM25公式
numerator = tf * (self.k1 + 1)
denominator = tf + self.k1 * (1 - self.b + self.b * (doc_len / self.avg_doc_len))
score += idf * (numerator / denominator)
return score
def batch_search(self, queries: List[str], top_k: int = 10) -> List[List[Dict[str, Any]]]:
"""批量搜索"""
results = []
for query in queries:
results.append(self.search(query, top_k))
return results
class VectorRetriever:
"""
向量检索器
基于BGE等向量模型进行语义检索
"""
def __init__(self, embedding_model: Any):
"""
初始化向量检索器
Args:
embedding_model: 向量化模型
"""
self.embedding_model = embedding_model
self.documents: List[Document] = []
self.doc_embeddings: Optional[np.ndarray] = None
self.is_built = False
def add_documents(self, documents: List[Dict[str, Any]]):
"""
添加文档到索引
Args:
documents: 文档列表
"""
for doc in documents:
doc_id = doc.get('doc_id', f"doc_{len(self.documents)}")
content = doc.get('content', '')
embedding = doc.get('embedding')
document = Document(
doc_id=doc_id,
content=content,
embedding=embedding,
metadata=doc.get('metadata', {})
)
self.documents.append(document)
# 构建向量索引
self._build_index()
logger.info(f"向量索引已构建,包含 {len(self.documents)} 个文档")
def _build_index(self):
"""构建向量索引"""
if not self.documents:
return
# 提取需要编码的文档
need_encode = []
encode_indices = []
for idx, doc in enumerate(self.documents):
if doc.embedding is None:
need_encode.append(doc.content)
encode_indices.append(idx)
# 编码文档
if need_encode:
embeddings = self.embedding_model.encode_documents(need_encode)
for idx, emb in zip(encode_indices, embeddings):
self.documents[idx].embedding = emb
# 构建向量矩阵
self.doc_embeddings = np.array([doc.embedding for doc in self.documents])
self.is_built = True
def search(self, query: str, top_k: int = 10) -> List[Dict[str, Any]]:
"""
执行向量搜索
Args:
query: 查询文本
top_k: 返回结果数量
Returns:
搜索结果列表
"""
if not self.is_built or self.doc_embeddings is None:
logger.warning("向量索引未构建")
return []
# 编码查询
query_embedding = self.embedding_model.encode_queries([query])
# 计算相似度
similarities = cosine_similarity(query_embedding, self.doc_embeddings)[0]
# 获取top-k索引
top_indices = np.argsort(similarities)[::-1][:top_k]
# 构建结果
results = []
for idx in top_indices:
doc = self.documents[idx]
results.append({
'doc_id': doc.doc_id,
'content': doc.content,
'score': float(similarities[idx]),
'index': idx,
'metadata': doc.metadata
})
return results
def batch_search(self, queries: List[str], top_k: int = 10) -> List[List[Dict[str, Any]]]:
"""批量搜索"""
results = []
for query in queries:
results.append(self.search(query, top_k))
return results
class HybridRetriever:
"""
混合检索器
结合向量检索和关键词检索,提高召回率
"""
def __init__(
self,
embedding_model: Any,
tokenizer: Optional[ChineseTokenizer] = None,
weight_vector: float = 0.5,
weight_keyword: float = 0.5,
use_normalization: bool = True
):
"""
初始化混合检索器
Args:
embedding_model: 向量化模型
tokenizer: 分词器
weight_vector: 向量检索权重
weight_keyword: 关键词检索权重
use_normalization: 是否归一化得分
"""
self.embedding_model = embedding_model
self.tokenizer = tokenizer or ChineseTokenizer()
self.weight_vector = weight_vector
self.weight_keyword = weight_keyword
self.use_normalization = use_normalization
# 初始化子检索器
self.vector_retriever = VectorRetriever(embedding_model)
self.bm25_retriever = BM25Retriever(tokenizer)
self.documents: List[Document] = []
self.is_built = False
def add_documents(self, documents: List[Dict[str, Any]]):
"""
添加文档到索引
Args:
documents: 文档列表
"""
# 为文档生成embedding(如果没有)
for doc in documents:
if 'embedding' not in doc or doc['embedding'] is None:
content = doc.get('content', '')
if content:
doc['embedding'] = self.embedding_model.encode([content])[0]
# 添加到向量检索器
self.vector_retriever.add_documents(documents)
# 添加到BM25检索器
self.bm25_retriever.add_documents(documents)
# 存储文档
for doc in documents:
doc_id = doc.get('doc_id', f"doc_{len(self.documents)}")
content = doc.get('content', '')
document = Document(
doc_id=doc_id,
content=content,
embedding=doc.get('embedding'),
metadata=doc.get('metadata', {})
)
self.documents.append(document)
self.is_built = True
logger.info(f"混合检索器已构建,包含 {len(self.documents)} 个文档")
def search(
self,
query: str,
top_k: int = 10,
vector_weight: Optional[float] = None,
keyword_weight: Optional[float] = None,
rerank: bool = True,
use_entropy_weight: bool = False
) -> List[Dict[str, Any]]:
"""
执行混合搜索
Args:
query: 查询文本
top_k: 返回结果数量
vector_weight: 向量检索权重(覆盖默认值)
keyword_weight: 关键词检索权重(覆盖默认值)
rerank: 是否使用重排序
use_entropy_weight: 是否使用熵权法动态调整权重
Returns:
搜索结果列表
"""
if not self.is_built:
logger.warning("混合检索器未构建")
return []
# 设置权重
w_vector = vector_weight if vector_weight is not None else self.weight_vector
w_keyword = keyword_weight if keyword_weight is not None else self.weight_keyword
# 执行向量检索
vector_results = self.vector_retriever.search(query, top_k=top_k * 2)
# 执行关键词检索
keyword_results = self.bm25_retriever.search(query, top_k=top_k * 2)
# 如果使用熵权法,动态计算权重
if use_entropy_weight:
w_vector, w_keyword = self._calculate_entropy_weights(vector_results, keyword_results)
logger.info(f"熵权法计算权重: vector={w_vector:.3f}, keyword={w_keyword:.3f}")
# 融合结果
fused_results = self._fuse_results(
vector_results,
keyword_results,
w_vector,
w_keyword
)
# 重排序(可选)
if rerank:
fused_results = self._rerank_results(fused_results, query)
# 返回top-k
return fused_results[:top_k]
def _fuse_results(
self,
vector_results: List[Dict[str, Any]],
keyword_results: List[Dict[str, Any]],
weight_vector: float,
weight_keyword: float
) -> List[Dict[str, Any]]:
"""
融合两种检索结果
Args:
vector_results: 向量检索结果
keyword_results: 关键词检索结果
weight_vector: 向量检索权重
weight_keyword: 关键词检索权重
Returns:
融合后的结果列表
"""
# 构建文档ID到得分的映射
doc_scores = defaultdict(float)
doc_info = {}
# 归一化向量检索得分
if self.use_normalization and vector_results:
max_score = max(r['score'] for r in vector_results)
min_score = min(r['score'] for r in vector_results)
score_range = max_score - min_score if max_score != min_score else 1.0
for result in vector_results:
doc_id = result['doc_id']
normalized_score = (result['score'] - min_score) / score_range
doc_scores[doc_id] += normalized_score * weight_vector
doc_info[doc_id] = {
'content': result['content'],
'metadata': result.get('metadata', {}),
'vector_score': result['score'],
'normalized_vector_score': normalized_score
}
else:
for result in vector_results:
doc_id = result['doc_id']
doc_scores[doc_id] += result['score'] * weight_vector
doc_info[doc_id] = {
'content': result['content'],
'metadata': result.get('metadata', {}),
'vector_score': result['score']
}
# 归一化关键词检索得分
if self.use_normalization and keyword_results:
max_score = max(r['score'] for r in keyword_results) if keyword_results else 1.0
min_score = min(r['score'] for r in keyword_results) if keyword_results else 0.0
score_range = max_score - min_score if max_score != min_score else 1.0
for result in keyword_results:
doc_id = result['doc_id']
normalized_score = (result['score'] - min_score) / score_range
doc_scores[doc_id] += normalized_score * weight_keyword
if doc_id in doc_info:
doc_info[doc_id]['keyword_score'] = result['score']
doc_info[doc_id]['normalized_keyword_score'] = normalized_score
else:
doc_info[doc_id] = {
'content': result['content'],
'metadata': result.get('metadata', {}),
'keyword_score': result['score'],
'normalized_keyword_score': normalized_score
}
else:
for result in keyword_results:
doc_id = result['doc_id']
doc_scores[doc_id] += result['score'] * weight_keyword
if doc_id not in doc_info:
doc_info[doc_id] = {
'content': result['content'],
'metadata': result.get('metadata', {}),
'keyword_score': result['score']
}
# 构建最终结果
final_results = []
for doc_id, total_score in doc_scores.items():
info = doc_info.get(doc_id, {})
final_results.append({
'doc_id': doc_id,
'content': info.get('content', ''),
'score': total_score,
'metadata': info.get('metadata', {}),
'details': {
'vector_score': info.get('vector_score', 0),
'keyword_score': info.get('keyword_score', 0),
'normalized_vector_score': info.get('normalized_vector_score', 0),
'normalized_keyword_score': info.get('normalized_keyword_score', 0)
}
})
# 按综合得分排序
final_results.sort(key=lambda x: x['score'], reverse=True)
return final_results
def _rerank_results(
self,
results: List[Dict[str, Any]],
query: str,
top_n: int = 20
) -> List[Dict[str, Any]]:
"""
重排序结果(使用交叉编码器或其它方法)
Args:
results: 待重排序的结果
query: 查询文本
top_n: 重排序的文档数量
Returns:
重排序后的结果
"""
# 这里可以集成交叉编码器进行重排序
# 简化版:基于查询和文档的字符重叠度进行微调
query_tokens = set(self.tokenizer.tokenize(query))
for result in results[:top_n]:
content_tokens = set(self.tokenizer.tokenize(result['content']))
overlap = len(query_tokens & content_tokens) / (len(query_tokens) + 0.1)
# 微调得分:重叠度越高,得分越高
result['score'] = result['score'] * (1 + 0.1 * overlap)
# 重新排序
results.sort(key=lambda x: x['score'], reverse=True)
return results
def _calculate_entropy_weights(
self,
vector_results: List[Dict[str, Any]],
keyword_results: List[Dict[str, Any]]
) -> Tuple[float, float]:
"""
使用熵权法计算权重
Args:
vector_results: 向量检索结果
keyword_results: 关键词检索结果
Returns:
(向量权重, 关键词权重)
"""
# 提取得分
vector_scores = np.array([r['score'] for r in vector_results]) if vector_results else np.array([])
keyword_scores = np.array([r['score'] for r in keyword_results]) if keyword_results else np.array([])
if len(vector_scores) == 0 or len(keyword_scores) == 0:
return 0.5, 0.5
# 归一化
max_v = np.max(vector_scores) if len(vector_scores) > 0 else 1
min_v = np.min(vector_scores) if len(vector_scores) > 0 else 0
range_v = max_v - min_v if max_v != min_v else 1
normalized_v = (vector_scores - min_v) / range_v
max_k = np.max(keyword_scores) if len(keyword_scores) > 0 else 1
min_k = np.min(keyword_scores) if len(keyword_scores) > 0 else 0
range_k = max_k - min_k if max_k != min_k else 1
normalized_k = (keyword_scores - min_k) / range_k
# 计算熵
def compute_entropy(scores):
if len(scores) == 0 or np.sum(scores) == 0:
return 1.0
p = scores / np.sum(scores)
p = np.clip(p, 1e-10, 1.0)
return -np.sum(p * np.log(p)) / np.log(len(scores))
entropy_v = compute_entropy(normalized_v)
entropy_k = compute_entropy(normalized_k)
# 计算权重
total_entropy = entropy_v + entropy_k
if total_entropy == 0:
return 0.5, 0.5
weight_v = (1 - entropy_v) / (2 - total_entropy)
weight_k = (1 - entropy_k) / (2 - total_entropy)
# 归一化
total_weight = weight_v + weight_k
if total_weight > 0:
weight_v /= total_weight
weight_k /= total_weight
return weight_v, weight_k
def batch_search(
self,
queries: List[str],
top_k: int = 10,
**kwargs
) -> List[List[Dict[str, Any]]]:
"""批量搜索"""
results = []
for query in queries:
results.append(self.search(query, top_k, **kwargs))
return results
class HybridSearchEngine:
"""
混合搜索引擎
提供完整的搜索功能,包括索引管理、搜索和评估
"""
def __init__(self, embedding_model: Any):
self.retriever = HybridRetriever(embedding_model)
self.search_history = []
def index_documents(self, documents: List[Dict[str, Any]]):
"""索引文档"""
self.retriever.add_documents(documents)
logger.info(f"搜索引擎索引完成,共索引 {len(documents)} 个文档")
def search(
self,
query: str,
top_k: int = 10,
**kwargs
) -> List[Dict[str, Any]]:
"""搜索"""
results = self.retriever.search(query, top_k, **kwargs)
# 记录搜索历史
self.search_history.append({
'query': query,
'results_count': len(results),
'timestamp': datetime.now().isoformat()
})
return results
def get_stats(self) -> Dict[str, Any]:
"""获取统计信息"""
return {
'total_documents': len(self.retriever.documents),
'total_searches': len(self.search_history),
'retriever_built': self.retriever.is_built
}
# 测试代码
def test_hybrid_retriever():
"""测试混合检索器"""
# 模拟BGE模型
class MockBGE:
def encode(self, texts):
if isinstance(texts, str):
texts = [texts]
return np.random.randn(len(texts), 768)
def encode_queries(self, texts):
return self.encode(texts)
def encode_documents(self, texts):
return self.encode(texts)
# 创建测试文档
test_documents = [
{
'doc_id': 'doc_001',
'content': '装备维修手册:发动机故障排查方法包括检查油路、电路和机械部件。'
},
{
'doc_id': 'doc_002',
'content': '向量检索技术用于快速找到相关维修案例和历史记录。'
},
{
'doc_id': 'doc_003',
'content': 'BM25算法是经典的信息检索方法,用于关键词匹配搜索。'
},
{
'doc_id': 'doc_004',
'content': '混合检索结合了语义理解和关键词匹配的优势,提高召回率。'
},
{
'doc_id': 'doc_005',
'content': '装备维修保障系统支持多种检索方式,包括向量检索和关键词检索。'
}
]
# 初始化混合检索器
embedding_model = MockBGE()
hybrid_retriever = HybridRetriever(
embedding_model=embedding_model,
weight_vector=0.6,
weight_keyword=0.4
)
# 添加文档
hybrid_retriever.add_documents(test_documents)
# 测试搜索
queries = [
'发动机故障排查',
'信息检索算法',
'维修保障系统'
]
for query in queries:
print(f"\n{'='*60}")
print(f"查询: {query}")
print(f"{'='*60}")
# 执行混合搜索
results = hybrid_retriever.search(query, top_k=3)
for i, result in enumerate(results, 1):
print(f"\n{i}. 文档ID: {result['doc_id']}")
print(f" 综合得分: {result['score']:.4f}")
print(f" 内容: {result['content']}")
print(f" 详情: {result.get('details', {})}")
return hybrid_retriever
if __name__ == "__main__":
import datetime
print("="*60)
print("混合检索算法测试")
print("="*60)
test_hybrid_retriever()
&spm=1001.2101.3001.5002&articleId=163749570&d=1&t=3&u=739e4b1b3a7b47759805f4b0fbb42269)
621

被折叠的 条评论
为什么被折叠?



