LangChain智能文档助手【5】-检索器
·
这是一个基于 LangChain 框架和通义千问(Qwen)大语言模型构建的检索增强生成(RAG)智能问答系统的教程。该系统是基于最基础的功能点,如大语言模型调用,文档处理,嵌入模型向量化,向量存储和检索,智能对话构建而成,最后用streamlit生成一个web界面,将功能可视化。
第五章 检索器
现了基于 Qwen 的智能检索器,是 RAG 系统的核心组件,按如下层级创建文件retriever.py。

retriever.py文件的完整代码:
import os
from typing import List, Dict, Any, Optional, Tuple
from dotenv import load_dotenv
from qwen_client import QwenClient
from vector_store import QwenVectorStore, DashScopeEmbeddings
from langchain_community.vectorstores import FAISS
from langchain_core.documents import Document
from langchain_core.retrievers import BaseRetriever
from langchain_core.callbacks import CallbackManagerForRetrieverRun
from pydantic import BaseModel, Field
load_dotenv()
# ==================== 自定义检索器类 ====================
class QwenQueryGenerator(BaseModel):
"""Qwen查询生成器"""
qwen_client: Any = Field(..., description="Qwen客户端")
num_queries: int = Field(default=3, description="生成的查询数量")
def generate(self, original_query: str) -> List[str]:
"""使用Qwen生成多个相关问题"""
prompt = f"""针对以下查询,生成{self.num_queries}个不同角度或表述的相关问题:
原始查询:{original_query}
要求:
1. 从不同角度提出问题
2. 使用不同的表述方式
3. 确保问题都与原始查询相关
4. 返回格式:每个问题单独一行
生成的问题:"""
try:
response = self.qwen_client.chat_completion([
{"role": "user", "content": prompt}
])
# 解析响应,每行一个问题
queries = [line.strip() for line in response.split('\n') if line.strip()]
queries = [q for q in queries if not q.startswith(('1.', '2.', '3.', '-', '•'))]
# 确保包含原始查询
if original_query not in queries:
queries.insert(0, original_query)
print(f"📝 生成的查询:{queries}")
return queries
except Exception as e:
print(f"❌ 生成多查询失败: {e}")
return [original_query]
class MultiQueryRetriever(BaseRetriever):
"""多查询检索器"""
base_retriever: BaseRetriever
query_generator: QwenQueryGenerator
def _get_relevant_documents(
self,
query: str,
*,
run_manager: CallbackManagerForRetrieverRun = None
) -> List[Document]:
"""执行多查询检索"""
# 生成多个查询
queries = self.query_generator.generate(query)
# 收集所有结果
all_documents = []
seen_content = set() # 用于去重
for q in queries:
try:
# 使用 invoke 方法获取文档
if hasattr(self.base_retriever, 'invoke'):
documents = self.base_retriever.invoke(q)
elif hasattr(self.base_retriever, '_get_relevant_documents'):
documents = self.base_retriever._get_relevant_documents(q)
else:
# 尝试直接调用
documents = self.base_retriever(q)
for doc in documents:
# 基于内容去重
content_hash = hash(doc.page_content[:500]) # 只考虑前500字符
if content_hash not in seen_content:
seen_content.add(content_hash)
all_documents.append(doc)
except Exception as e:
print(f"⚠️ 查询 '{q}' 检索失败: {e}")
continue
# 返回结果
return all_documents
def invoke(self, query: str) -> List[Document]:
"""兼容 invoke 方法"""
return self._get_relevant_documents(query)
class EnsembleRetriever(BaseRetriever):
"""混合检索器"""
retrievers: List[BaseRetriever] = Field(..., description="多个检索器")
weights: List[float] = Field(default=None, description="各检索器的权重")
def __init__(self, retrievers: List[BaseRetriever], weights: Optional[List[float]] = None, **kwargs):
if weights is None:
weights = [1.0 / len(retrievers)] * len(retrievers)
if len(retrievers) != len(weights):
raise ValueError("检索器和权重数量必须一致")
super().__init__(retrievers=retrievers, weights=weights, **kwargs)
def _get_relevant_documents(
self,
query: str,
*,
run_manager: CallbackManagerForRetrieverRun = None
) -> List[Document]:
"""执行混合检索"""
all_documents = []
for retriever, weight in zip(self.retrievers, self.weights):
try:
# 使用 invoke 方法获取文档
if hasattr(retriever, 'invoke'):
documents = retriever.invoke(query)
elif hasattr(retriever, '_get_relevant_documents'):
documents = retriever._get_relevant_documents(query)
else:
# 尝试直接调用
documents = retriever(query)
# 为每个文档添加权重
for doc in documents:
doc.metadata = doc.metadata or {}
doc.metadata["retriever_weight"] = weight
doc.metadata["retriever_type"] = type(retriever).__name__
all_documents.extend(documents)
except Exception as e:
print(f"⚠️ 检索器 {type(retriever).__name__} 失败: {e}")
continue
# 去重并加权排序
return self._deduplicate_and_rank(all_documents)
def _deduplicate_and_rank(self, documents: List[Document]) -> List[Document]:
"""去重并加权排序"""
seen = {}
for doc in documents:
# 使用内容作为唯一标识
content_key = doc.page_content[:500] # 前500字符
if content_key in seen:
# 合并权重
seen[content_key]["weight"] += doc.metadata.get("retriever_weight", 0)
else:
seen[content_key] = {
"doc": doc,
"weight": doc.metadata.get("retriever_weight", 0)
}
# 按权重排序
sorted_docs = sorted(seen.values(), key=lambda x: x["weight"], reverse=True)
return [item["doc"] for item in sorted_docs]
def invoke(self, query: str) -> List[Document]:
"""兼容 invoke 方法"""
return self._get_relevant_documents(query)
class QwenRerankRetriever(BaseRetriever):
"""Qwen重排序检索器"""
base_retriever: BaseRetriever
qwen_client: Any
k: int = Field(default=8, description="初始检索数量")
rerank_k: int = Field(default=3, description="重排序后数量")
def _get_relevant_documents(
self,
query: str,
*,
run_manager: CallbackManagerForRetrieverRun = None
) -> List[Document]:
"""执行重排序检索"""
# 第一步:基础检索
if hasattr(self.base_retriever, 'invoke'):
docs = self.base_retriever.invoke(query)
elif hasattr(self.base_retriever, '_get_relevant_documents'):
docs = self.base_retriever._get_relevant_documents(query)
else:
docs = self.base_retriever(query)
if len(docs) <= self.rerank_k:
return docs[:self.rerank_k]
# 第二步:使用Qwen进行重排序
try:
# 构建重排序提示
docs_text = "\n\n".join([
f"文档{i+1}:\n"
f"内容: {doc.page_content[:200]}...\n"
f"元数据: {doc.metadata if hasattr(doc, 'metadata') and doc.metadata else '无'}"
for i, doc in enumerate(docs[:self.k])
])
prompt = f"""请根据以下标准为文档评分并排序(1-5分,5分最相关):
查询:{query}
评分标准:
1. 与查询的直接相关性(40%)
2. 信息的完整性和有用性(30%)
3. 时效性(根据元数据判断,20%)
4. 来源权威性(10%)
文档列表:
{docs_text}
请按以下格式返回排序结果:
文档编号:总分数(相关分,完整分,时效分,权威分)
例如:3:4.5(5,4,4,5), 1:4.2(4,5,3,5), 2:3.8(4,3,4,4)"""
response = self.qwen_client.chat_completion([
{"role": "user", "content": prompt}
])
print(f"🤖 Qwen重排序响应: {response}")
# 解析排序结果
ranked_docs = self._parse_ranking(response, docs)
return ranked_docs[:self.rerank_k]
except Exception as e:
print(f"⚠️ 重排序失败,返回原始结果: {e}")
return docs[:self.rerank_k]
def _parse_ranking(self, response: str, docs: List[Document]) -> List[Document]:
"""解析排序结果"""
try:
rankings = []
response = response.strip()
# 尝试解析复杂格式 "3:4.5(5,4,4,5), 1:4.2(4,5,3,5)"
for item in response.split(','):
item = item.strip()
if ':' in item:
# 提取文档编号
if '(' in item:
doc_part = item.split('(')[0]
else:
doc_part = item
if ':' in doc_part:
try:
doc_idx_str, _ = doc_part.split(':')
doc_idx = int(doc_idx_str.strip()) - 1
if 0 <= doc_idx < len(docs):
rankings.append((doc_idx, 1.0))
except ValueError:
continue
if rankings:
# 按文档在排序列表中的顺序返回
return [docs[idx] for idx, _ in rankings if idx < len(docs)]
else:
# 尝试简单格式
for i, char in enumerate(response):
if char.isdigit():
try:
doc_idx = int(char) - 1
if 0 <= doc_idx < len(docs):
rankings.append((doc_idx, 1.0))
except:
pass
if rankings:
return [docs[idx] for idx, _ in rankings]
except Exception as e:
print(f"⚠️ 排序解析失败: {e}")
# 默认返回原始顺序
return docs
def invoke(self, query: str) -> List[Document]:
"""兼容 invoke 方法"""
return self._get_relevant_documents(query)
# ==================== 辅助函数 ====================
def get_documents_from_retriever(retriever, query: str) -> List[Document]:
"""从检索器获取文档,兼容不同接口"""
if hasattr(retriever, 'invoke'):
return retriever.invoke(query)
elif hasattr(retriever, '_get_relevant_documents'):
return retriever._get_relevant_documents(query)
else:
# 尝试直接调用
return retriever(query)
# ==================== 主检索器类 ====================
class QwenRetriever:
"""Qwen优化的智能检索器"""
def __init__(self, vectorstore_path: str = "./faiss_db_qwen"):
self.vectorstore_path = vectorstore_path
self.qwen_client = QwenClient()
self.vectorstore = None
self.retriever = None
# 初始化向量存储管理器
self.vector_manager = QwenVectorStore(persist_dir=vectorstore_path)
# 初始化向量存储
self._init_vectorstore()
def _init_vectorstore(self):
"""初始化向量存储"""
try:
# 尝试加载现有的向量存储
self.vectorstore = self.vector_manager.load_existing_store()
if self.vectorstore:
count = len(self.vectorstore.index_to_docstore_id)
print(f"✅ FAISS向量存储加载成功,包含 {count} 个文档")
else:
print("⚠️ 未找到现有向量存储,请先创建或添加文档")
except Exception as e:
print(f"❌ 加载向量存储失败: {e}")
def create_basic_retriever(self, search_type: str = "similarity", k: int = 4):
"""创建基础检索器"""
if not self.vectorstore:
print("❌ 向量存储未初始化")
return None
# FAISS的检索器配置
search_kwargs = {"k": k}
if search_type == "similarity":
# 相似度搜索
self.retriever = self.vectorstore.as_retriever(
search_type="similarity",
search_kwargs=search_kwargs
)
print(f"✅ 创建相似度检索器,返回前{k}个最相关结果")
elif search_type == "mmr":
# 最大边际相关性搜索
search_kwargs["fetch_k"] = min(k * 3, 20)
self.retriever = self.vectorstore.as_retriever(
search_type="mmr",
search_kwargs=search_kwargs
)
print(f"✅ 创建MMR检索器,返回{k}个多样化结果")
else:
print(f"❌ 不支持的搜索类型: {search_type}")
return None
return self.retriever
def create_multi_query_retriever(self, k: int = 4, num_queries: int = 3):
"""创建多查询检索器"""
if not self.vectorstore:
print("❌ 向量存储未初始化")
return None
try:
# 创建基础检索器
base_retriever = self.vectorstore.as_retriever(
search_kwargs={"k": k}
)
# 创建查询生成器
query_generator = QwenQueryGenerator(
qwen_client=self.qwen_client,
num_queries=num_queries
)
# 创建多查询检索器
self.retriever = MultiQueryRetriever(
base_retriever=base_retriever,
query_generator=query_generator
)
print(f"✅ 创建多查询检索器")
print(f" 生成{num_queries}个相关查询,每个查询返回{k}个结果")
return self.retriever
except Exception as e:
print(f"❌ 创建多查询检索器失败: {e}")
# 降级为使用基础检索器
return self.create_basic_retriever(k=k)
def create_hybrid_retriever(self, k: int = 3, weights: Optional[List[float]] = None):
"""创建混合检索器"""
if not self.vectorstore:
print("❌ 向量存储未初始化")
return None
try:
# 创建不同的检索器
retrievers = []
# 1. 相似度检索器
similarity_retriever = self.vectorstore.as_retriever(
search_type="similarity",
search_kwargs={"k": k * 2}
)
retrievers.append(similarity_retriever)
# 2. MMR检索器
mmr_retriever = self.vectorstore.as_retriever(
search_type="mmr",
search_kwargs={"k": k * 2, "fetch_k": k * 4}
)
retrievers.append(mmr_retriever)
# 3. 带分数的检索器(自定义)
class ScoreBasedRetriever(BaseRetriever):
def __init__(self, vectorstore, k):
self.vectorstore = vectorstore
self.k = k
def _get_relevant_documents(self, query: str, **kwargs) -> List[Document]:
# 使用带分数的搜索
results = self.vectorstore.similarity_search_with_score(query, k=self.k)
documents = []
for doc, score in results:
doc.metadata = doc.metadata or {}
doc.metadata["similarity_score"] = score
documents.append(doc)
return documents
def invoke(self, query: str) -> List[Document]:
return self._get_relevant_documents(query)
score_retriever = ScoreBasedRetriever(self.vectorstore, k=k*2)
retrievers.append(score_retriever)
# 设置权重
if weights is None:
weights = [0.4, 0.3, 0.3] # 对应上面的三个检索器
# 创建混合检索器
self.retriever = EnsembleRetriever(
retrievers=retrievers,
weights=weights
)
print(f"✅ 创建混合检索器")
print(f" 结合{len(retrievers)}种检索策略")
print(f" 权重分配: {weights}")
return self.retriever
except Exception as e:
print(f"❌ 创建混合检索器失败: {e}")
# 降级为MMR检索器
return self.create_basic_retriever(search_type="mmr", k=k)
def create_qwen_rerank_retriever(self, k: int = 8, rerank_k: int = 3):
"""创建Qwen重排序检索器"""
if not self.vectorstore:
print("❌ 向量存储未初始化")
return None
try:
# 创建基础检索器
base_retriever = self.vectorstore.as_retriever(
search_kwargs={"k": k}
)
# 创建重排序检索器
self.retriever = QwenRerankRetriever(
base_retriever=base_retriever,
qwen_client=self.qwen_client,
k=k,
rerank_k=rerank_k
)
print(f"✅ 创建Qwen重排序检索器")
print(f" 先检索{k}个,再用Qwen重排序选择{rerank_k}个")
return self.retriever
except Exception as e:
print(f"❌ 创建重排序检索器失败: {e}")
return self.create_basic_retriever(k=rerank_k)
def create_compression_retriever(self, k: int = 4):
"""创建文档压缩检索器"""
if not self.vectorstore:
print("❌ 向量存储未初始化")
return None
try:
# 创建基础检索器(检索更多文档)
base_retriever = self.vectorstore.as_retriever(
search_kwargs={"k": k * 3}
)
# 自定义压缩检索器
class CompressionRetriever(BaseRetriever):
def __init__(self, base_retriever, qwen_client, k):
self.base_retriever = base_retriever
self.qwen_client = qwen_client
self.k = k
def _get_relevant_documents(self, query: str, **kwargs) -> List[Document]:
# 先获取多个文档
if hasattr(self.base_retriever, 'invoke'):
docs = self.base_retriever.invoke(query)
else:
docs = self.base_retriever._get_relevant_documents(query)
compressed_docs = []
for doc in docs[:self.k * 2]: # 处理前2k个文档
try:
# 构建压缩提示
prompt = f"""请从以下文档中提取与查询最相关的部分:
查询:{query}
文档内容:
{doc.page_content[:500]}
要求:
1. 只提取与查询直接相关的内容
2. 保持语义完整
3. 如果完全不相关,返回"空"
提取的内容:"""
response = self.qwen_client.chat_completion([
{"role": "user", "content": prompt}
])
if response and response.strip() and response.strip() != "空":
compressed_doc = Document(
page_content=response.strip(),
metadata=doc.metadata
)
compressed_docs.append(compressed_doc)
except Exception as e:
print(f"⚠️ 文档压缩失败: {e}")
# 保留原始文档作为降级方案
compressed_docs.append(doc)
return compressed_docs[:self.k]
def invoke(self, query: str) -> List[Document]:
return self._get_relevant_documents(query)
# 创建压缩检索器
self.retriever = CompressionRetriever(
base_retriever=base_retriever,
qwen_client=self.qwen_client,
k=k
)
print(f"✅ 创建文档压缩检索器")
print(f" 先检索{k*3}个,再用Qwen压缩提取{k}个最相关部分")
return self.retriever
except Exception as e:
print(f"❌ 创建压缩检索器失败: {e}")
return self.create_basic_retriever(k=k)
def add_documents(self, documents: List[Document]):
"""向向量库添加文档"""
return self.vector_manager.create_vector_store(documents)
def test_retrieval(self, query: str, retriever_type: str = "basic", **kwargs):
"""测试检索效果"""
print(f"\n🧪 测试检索器: {retriever_type}")
print(f"📝 查询: {query}")
# 创建对应的检索器
if retriever_type == "basic":
search_type = kwargs.get("search_type", "similarity")
k = kwargs.get("k", 4)
self.create_basic_retriever(search_type=search_type, k=k)
elif retriever_type == "multi_query":
k = kwargs.get("k", 4)
num_queries = kwargs.get("num_queries", 3)
self.create_multi_query_retriever(k=k, num_queries=num_queries)
elif retriever_type == "hybrid":
k = kwargs.get("k", 3)
weights = kwargs.get("weights", None)
self.create_hybrid_retriever(k=k, weights=weights)
elif retriever_type == "rerank":
k = kwargs.get("k", 8)
rerank_k = kwargs.get("rerank_k", 3)
self.create_qwen_rerank_retriever(k=k, rerank_k=rerank_k)
elif retriever_type == "compression":
k = kwargs.get("k", 4)
self.create_compression_retriever(k=k)
else:
print(f"❌ 不支持的检索器类型: {retriever_type}")
return
if not self.retriever:
print("❌ 检索器创建失败")
return
try:
# 执行检索 - 使用统一的函数
results = get_documents_from_retriever(self.retriever, query)
print(f"📊 检索到 {len(results)} 个结果:")
for i, doc in enumerate(results):
print(f"\n--- 结果 {i+1} ---")
print(f"内容: {doc.page_content[:200]}...")
if hasattr(doc, 'metadata') and doc.metadata:
print(f"元数据: {doc.metadata}")
return results
except Exception as e:
print(f"❌ 检索失败: {e}")
return None
# 使用示例
if __name__ == "__main__":
# 创建检索器
retriever = QwenRetriever()
# 测试基础检索
print("\n=== 测试基础检索 ===")
retriever.test_retrieval("什么是机器学习?", "basic")
# 测试多查询检索
print("\n=== 测试多查询检索 ===")
retriever.test_retrieval("人工智能的应用", "multi_query")
# 测试混合检索
print("\n=== 测试混合检索 ===")
retriever.test_retrieval("深度学习框架", "hybrid")
# 测试重排序检索
print("\n=== 测试重排序检索 ===")
retriever.test_retrieval("自然语言处理", "rerank", k=6, rerank_k=2)
运行代码
python src/retriever.py
代码解析:
- QwenRetriever - 工厂类/主控制器
# 使用方式
retriever = QwenRetriever() # ① 创建助理
results = retriever.test_retrieval("什么是AI?", "basic") # ② 问问题
功能:主控制器,帮你管理所有检索操作
2. 五种检索助手(由QwenRetriever管理)

方式1️⃣:快速查找(基础检索)create_basic_retriever
# 就像百度搜索,直接找最相似的
results = assistant.test_retrieval(
query="什么是机器学习?", # 你的问题
retriever_type="basic", # 选择基础检索
search_type="similarity", # 相似度搜索(默认)
k=3 # 返回3个结果
)
# 或者使用MMR搜索(避免重复内容)
results = assistant.test_retrieval(
query="人工智能技术",
retriever_type="basic",
search_type="mmr", # 多样性搜索
k=4 # 返回4个多样化的结果
)
方式2️⃣:多角度查找(多查询检索)create_multi_query_retriever
# 系统会帮你生成3个相关问题一起搜索
results = assistant.test_retrieval(
query="人工智能的应用",
retriever_type="multi_query", # 多查询检索
num_queries=3, # 生成3个相关问题
k=4 # 返回4个结果
)
# 原理:你问1个,系统帮你问多个
# 你: "人工智能的应用"
# 系统帮你问:
# 1. "人工智能在医疗领域的应用"
# 2. "人工智能在教育中的应用"
# 3. "人工智能技术实际使用场景"
方式3️⃣:智能推荐(重排序检索)create_qwen_rerank_retriever
# 先找8个候选,让AI专家挑最好的3个
results = assistant.test_retrieval(
query="深度学习框架比较",
retriever_type="rerank", # 重排序检索
k=8, # 先找8个候选
rerank_k=3 # AI推荐最好的3个
)
# AI专家评分标准:
# 1. 相关性(40%)
# 2. 信息完整性(30%)
# 3. 时效性(20%)
# 4. 权威性(10%)
方式4️⃣:综合决策(混合检索器)create_hybrid_retriever
# 结合多种策略,给出最平衡的结果
results = assistant.test_retrieval(
query="大数据分析技术",
retriever_type="hybrid", # 混合检索
k=3, # 返回3个结果
weights=[0.4, 0.3, 0.3] # 三种策略的权重
)
# 三种策略组合:
# 1. 相似度专家(权重40%):找最相似的
# 2. 多样性专家(权重30%):找不同类型的
# 3. 质量专家(权重30%):找质量高的
# 最终:三位专家投票决定
方式5️⃣:精华提取(压缩检索器)create_compression_retriever
# 只给你最核心的内容,节省阅读时间
results = assistant.test_retrieval(
query="Python异步编程",
retriever_type="compression", # 压缩检索
k=2 # 返回2个精炼结果
)
# 处理过程:
# 原文:"这是一篇关于Python异步编程的长文章...(500字)"
# ↓ 压缩 ↓
# 结果:"Python异步编程核心:async/await语法、事件循环、协程...(100字)"
决策流程图
开始 → 你的需求是什么?
↓
[要快] → 基础检索器(basic)
↓
[要准] → 重排序检索器(rerank)
↓
[要全] → 多查询检索器(multi_query)
↓
[要多样] → 混合检索器(hybrid)
↓
[要精简] → 压缩检索器(compression)
3. 方法的调用链示例
# 用户使用流程
retriever = QwenRetriever() # 1. 创建工厂
# 2. 选择需要的产品
multi_query_retriever = retriever.create_multi_query_retriever(
k=4,
num_queries=3
)
# 3. 使用产品
results = multi_query_retriever.invoke("什么是机器学习?")
# 内部发生了什么:
# 1. QwenRetriever.create_multi_query_retriever() 被调用
# 2. 创建了 QwenQueryGenerator 实例
# 3. 创建了 MultiQueryRetriever 实例
# 4. MultiQueryRetriever 在内部使用 QwenQueryGenerator
以上,我们完成了5种检索器的构建。接下来,我们把以上所学串联起来,构建一个简单的基于知识库的智能对话系统。
更多推荐



所有评论(0)