RAG系统的优化是一个系统工程,需要从数据预处理、检索策略、上下文增强、生成优化四个维度综合调优。下面我从诊断框架到具体优化策略全面讲解。
RAG系统的问题通常可以分为三类:
问题类型典型表现根本原因检索失败回答"找不到相关信息"检索结果不相关或遗漏生成失败回答与检索内容矛盾/忽略模型未正确使用上下文综合失败回答部分正确但遗漏关键点多跳推理或信息整合不足
python
复制
下载
# RAG系统评估维度
评估指标 = {
"检索质量": {
"召回率": "相关文档被检索到的比例",
"精确率": "检索结果中相关文档的比例",
"MRR": "第一个相关结果的排名",
"NDCG": "考虑排序的相关性质量"
},
"生成质量": {
"答案准确性": "事实正确性",
"完整性": "是否覆盖所有要点",
"引用准确性": "是否正确标注来源",
"幻觉率": "编造信息的比例"
},
"端到端": {
"延迟": "总响应时间",
"token消耗": "总输入+输出token数"
}
} text
复制
下载
1. 检索质量评估
├─ 手动检查Top-K结果相关性
├─ 分析失败案例
└─ 决策:检索优化 OR 生成优化
2. 生成质量评估
├─ 检查回答是否基于检索内容
├─ 检查是否有幻觉
└─ 决策:提示词优化 OR 模型微调
3. 端到端测试
├─ 建立测试集
├─ 自动化评估
└─ 持续迭代 策略参数优化方向块大小200-1000 tokens小块提升精确匹配,大块保留上下文重叠率10-20%避免边界信息切断分块策略语义/结构/固定根据文档类型选择
最佳实践:
python
复制
下载
from langchain.text_splitter import RecursiveCharacterTextSplitter
# 场景1:代码文档 - 保留函数完整性
code_splitter = RecursiveCharacterTextSplitter(
chunk_size=800,
chunk_overlap=100,
separators=["\nclass ", "\ndef ", "\n ", "\n", " ", ""]
)
# 场景2:Markdown文档 - 保留标题层级
markdown_splitter = MarkdownHeaderTextSplitter(
headers_to_split_on=[
("#", "h1"),
("##", "h2"),
("###", "h3")
]
)
# 场景3:长文档 - 父子文档索引
from langchain.retrievers import ParentDocumentRetriever
parent_splitter = RecursiveCharacterTextSplitter(chunk_size=2000)
child_splitter = RecursiveCharacterTextSplitter(chunk_size=400)
retriever = ParentDocumentRetriever(
vectorstore=vectorstore,
docstore=docstore,
child_splitter=child_splitter,
parent_splitter=parent_splitter
) python
复制
下载
# 为每个chunk添加丰富的元数据
def enrich_metadata(doc, chunk):
return {
"content": chunk,
"metadata": {
"source": doc.metadata.get("source"),
"page": doc.metadata.get("page"),
"section": extract_section(doc),
"title": extract_title(doc),
"chunk_index": chunk_index,
"total_chunks": total_chunks,
"keywords": extract_keywords(chunk), # 关键词提取
"entity_names": extract_entities(chunk), # 实体识别
"timestamp": doc.metadata.get("date"),
"author": doc.metadata.get("author"),
"embedding_model": "bge-large-zh-v1.5"
}
}
# 利用元数据过滤
results = vectorstore.similarity_search(
query,
filter={
"source": {"$in": ["policy.pdf", "guide.pdf"]},
"page": {"$gte": 10, "$lte": 50}
}
) 结合向量检索(语义)和关键词检索(精确),互补优势。
python
复制
下载
from langchain.retrievers import EnsembleRetriever
from langchain.retrievers import BM25Retriever
# 创建BM25检索器(关键词)
bm25_retriever = BM25Retriever.from_documents(documents)
bm25_retriever.k = 10
# 创建向量检索器
vector_retriever = vectorstore.as_retriever(search_kwargs={"k": 10})
# 融合检索器
ensemble_retriever = EnsembleRetriever(
retrievers=[bm25_retriever, vector_retriever],
weights=[0.3, 0.7] # 权重可调
)
# 执行
results = ensemble_retriever.invoke(query) 权重调优策略:
精确查询(产品型号、专有名词)→ 提高BM25权重
语义查询(概念解释、总结归纳)→ 提高向量权重
平衡场景 → 5:5 或 4:6
将用户问题优化为更适合检索的形式。
python
复制
下载
from langchain.chains import LLMChain
# 查询改写Prompt
query_rewrite_prompt = PromptTemplate(
template="""你是一个查询优化专家。将用户问题改写为更适合向量检索的形式。
原始问题:{question}
改写要求:
1. 提取核心关键词
2. 补充同义词和相关概念
3. 拆解复杂问题为多个简单问题
4. 保持原意不变
输出格式(JSON):
{{
"optimized_query": "优化后的主查询",
"sub_queries": ["子查询1", "子查询2"],
"keywords": ["关键词1", "关键词2"]
}}
""",
input_variables=["question"]
)
query_rewriter = LLMChain(llm=llm, prompt=query_rewrite_prompt)
# 使用
rewritten = query_rewriter.run(question="如何使用向量数据库优化RAG性能?")
# 输出:
# {
# "optimized_query": "向量数据库 RAG 性能优化 使用方法",
# "sub_queries": ["向量数据库选型", "RAG检索优化策略", "向量索引参数调优"],
# "keywords": ["向量数据库", "RAG", "性能优化", "索引", "检索"]
# }
# 多路召回
all_results = []
all_results.extend(vectorstore.search(rewritten["optimized_query"], k=5))
for sub_query in rewritten["sub_queries"]:
all_results.extend(vectorstore.search(sub_query, k=3)) 检索Top-K后,用更精确的模型重新排序。
python
复制
下载
from sentence_transformers import CrossEncoder
class Reranker:
def __init__(self, model_name="BAAI/bge-reranker-large"):
self.reranker = CrossEncoder(model_name)
def rerank(self, query, documents, top_k=5):
# 构建(查询,文档)对
pairs = [[query, doc.page_content] for doc in documents]
# 计算相关性分数
scores = self.reranker.predict(pairs)
# 按分数排序
ranked = sorted(
zip(documents, scores),
key=lambda x: x[1],
reverse=True
)
return [doc for doc, score in ranked[:top_k]]
# 使用
retriever = vectorstore.as_retriever(search_kwargs={"k": 20}) # 先召回20个
candidates = retriever.invoke(query)
reranker = Reranker()
top_docs = reranker.rerank(query, candidates, top_k=5) # 重排后取前5 重排序模型对比:
模型维度速度精度适用场景bge-reranker-large1024慢极高高精度需求bge-reranker-base768中高平衡场景cross-encoder/ms-marco768中高英文场景无重排序-快中实时性要求高
根据查询复杂度动态调整检索策略。
python
复制
下载
class AdaptiveRetriever:
def __init__(self, vectorstore, llm):
self.vectorstore = vectorstore
self.llm = llm
def retrieve(self, query):
# 1. 判断查询复杂度
complexity = self._assess_complexity(query)
# 2. 自适应策略
if complexity == "simple":
# 简单问题:少量精确检索
k = 3
use_rerank = False
elif complexity == "medium":
# 中等:标准检索
k = 5
use_rerank = True
else:
# 复杂:多路召回+重排序
k = 10
use_rerank = True
use_multi_query = True
# 3. 执行检索
return self._search(query, k, use_rerank, use_multi_query)
def _assess_complexity(self, query):
prompt = f"""判断问题复杂度(simple/medium/complex):
问题:{query}
判断标准:
- simple:事实查询,1-2个实体
- medium:需要简单推理,3-5个实体
- complex:多步推理,多个条件
只返回simple/medium/complex"""
return self.llm.invoke(prompt).content.strip() 检索结果过长时,进行压缩以减少噪音。
python
复制
下载
from langchain.retrievers import ContextualCompressionRetriever
from langchain.retrievers.document_compressors import LLMChainExtractor
# 创建压缩器
compressor = LLMChainExtractor.from_llm(llm)
# 包装检索器
compression_retriever = ContextualCompressionRetriever(
base_compressor=compressor,
base_retriever=vectorstore.as_retriever()
)
# 自动提取与问题最相关的片段
compressed_docs = compression_retriever.invoke(query) 检索句子级别,返回上下文窗口。
python
复制
下载
class SentenceWindowRetriever:
def __init__(self, vectorstore, window_size=2):
self.vectorstore = vectorstore
self.window_size = window_size # 前后各取N句
def retrieve(self, query, k=5):
# 1. 检索句子
sentences = self.vectorstore.similarity_search(query, k=k)
# 2. 扩展上下文窗口
expanded_docs = []
for sent in sentences:
sent_idx = sent.metadata["sentence_index"]
doc_id = sent.metadata["doc_id"]
# 获取前后N句
context = self._get_surrounding_sentences(
doc_id, sent_idx, self.window_size
)
expanded_docs.append(context)
return expanded_docs python
复制
下载
def build_context(retrieved_docs, max_tokens=2000):
"""动态构建上下文,不超过token限制"""
context = []
current_tokens = 0
# 按相关性分数排序
for doc in sorted(retrieved_docs, key=lambda x: x.score, reverse=True):
doc_tokens = num_tokens_from_string(doc.page_content)
if current_tokens + doc_tokens <= max_tokens:
context.append(doc.page_content)
current_tokens += doc_tokens
else:
# 截断最后一个文档
remaining = max_tokens - current_tokens
if remaining > 100: # 剩余足够多才截断
truncated = truncate_text(doc.page_content, remaining)
context.append(truncated)
break
return "\n\n".join(context) python
复制
下载
# 高质量RAG提示词模板
RAG_PROMPT = """你是一个专业的知识助手,基于提供的参考资料回答问题。
## 参考资料
{context}
## 回答要求
1. **准确性**:严格基于参考资料,不要编造信息
2. **完整性**:覆盖所有相关要点
3. **可追溯**:用[来源: X]标注信息来源
4. **结构化**:复杂问题使用分点或表格
5. **谦逊**:不确定时明确说明
## 特殊处理
- 如果资料不足:回答"基于现有资料,我找到以下信息:...,但这可能不完整"
- 如果资料冲突:说明冲突点,并标注各自来源
- 如果问题超出范围:礼貌说明并建议咨询方向
## 用户问题
{question}
## 回答
"""
# 带思维链的复杂推理
RAG_PROMPT_WITH_COT = """基于参考资料回答问题。请先分析,再回答。
## 参考资料
{context}
## 分析步骤
1. 提取关键信息:[从资料中找到的关键点]
2. 识别关系:[信息之间的关联]
3. 推理过程:[如何得出结论]
## 用户问题
{question}
## 最终答案
""" 一次检索不够时,基于已有信息继续检索。
python
复制
下载
class IterativeRetriever:
def __init__(self, retriever, llm, max_iterations=3):
self.retriever = retriever
self.llm = llm
self.max_iterations = max_iterations
def retrieve(self, query):
all_docs = []
current_query = query
for i in range(self.max_iterations):
# 检索
docs = self.retriever.invoke(current_query)
all_docs.extend(docs)
# 判断是否需要继续
if self._is_sufficient(docs, query):
break
# 生成新查询
current_query = self._generate_next_query(query, all_docs)
return self._deduplicate(all_docs)
def _generate_next_query(self, original_query, retrieved_docs):
prompt = f"""基于已检索到的信息,生成一个更精准的检索查询。
原始问题:{original_query}
已检索到的信息摘要:
{self._summarize(retrieved_docs)}
需要补充什么信息?生成一个查询来获取它。
只返回查询,不要其他内容。"""
return self.llm.invoke(prompt).content.strip() 生成答案后,自我检查并修正。
python
复制
下载
class SelfCorrectingChain:
def __init__(self, retriever, llm):
self.retriever = retriever
self.llm = llm
def invoke(self, query):
# 1. 初次生成
context = self.retriever.invoke(query)
initial_answer = self._generate(query, context)
# 2. 自我检查
issues = self._check_answer(query, context, initial_answer)
# 3. 如果有问题,修正
if issues:
corrected = self._correct(query, context, initial_answer, issues)
return corrected
return initial_answer
def _check_answer(self, query, context, answer):
prompt = f"""检查答案是否有问题:
问题:{query}
参考资料:{context}
答案:{answer}
检查项:
1. 答案是否基于参考资料?
2. 是否有事实错误?
3. 是否有遗漏要点?
4. 是否有幻觉信息?
如有问题,列出具体问题。如无问题,返回"无问题"。
"""
return self.llm.invoke(prompt).content.strip() 不同类型问题路由到不同处理流程。
python
复制
下载
class QueryRouter:
def __init__(self, routes):
self.routes = routes # {"类型": 处理链}
self.classifier = self._create_classifier()
def route(self, query):
# 1. 分类
category = self.classifier.classify(query)
# 2. 路由到对应处理链
if category == "factual":
return self.routes["factual"].invoke(query) # 简单检索
elif category == "comparison":
return self.routes["comparison"].invoke(query) # 多路检索+对比
elif category == "calculation":
return self.routes["calculation"].invoke(query) # 检索数据+计算
else:
return self.routes["general"].invoke(query) # 标准RAG
def _create_classifier(self):
# 使用小模型快速分类
return ClassificationChain.from_llm(
llm=ChatOpenAI(model="gpt-3.5-turbo"),
categories=["factual", "comparison", "calculation", "general"],
prompt=classification_prompt
) python
复制
下载
def fusion_retrieval(query, retrievers, weights=None, k=10):
"""
融合多个检索器的结果
使用RRF(Reciprocal Rank Fusion)算法
"""
if weights is None:
weights = [1.0] * len(retrievers)
# 收集所有结果
all_results = []
for retriever in retrievers:
results = retriever.invoke(query)
all_results.append(results)
# RRF融合
scores = {}
for idx, results in enumerate(all_results):
for rank, doc in enumerate(results):
doc_id = doc.metadata.get("id", doc.page_content)
# RRF公式: 1/(k + rank)
rrf_score = weights[idx] / (60 + rank + 1)
scores[doc_id] = scores.get(doc_id, 0) + rrf_score
# 按融合分数排序
sorted_docs = sorted(scores.items(), key=lambda x: x[1], reverse=True)
return sorted_docs[:k] python
复制
下载
from functools import lru_cache
import hashlib
class CachedRetriever:
def __init__(self, retriever, cache_ttl=3600):
self.retriever = retriever
self.cache = {} # 实际应用用Redis
self.cache_ttl = cache_ttl
def _get_cache_key(self, query):
return hashlib.md5(query.encode()).hexdigest()
def retrieve(self, query):
cache_key = self._get_cache_key(query)
# 检查缓存
if cache_key in self.cache:
cached = self.cache[cache_key]
if time.time() - cached["timestamp"] < self.cache_ttl:
return cached["results"]
# 实际检索
results = self.retriever.invoke(query)
# 更新缓存
self.cache[cache_key] = {
"results": results,
"timestamp": time.time()
}
return results 优化项实施难度效果提升优先级适用场景分块大小调优⭐⭐⭐⭐⭐高所有场景元数据增强⭐⭐⭐⭐⭐中结构化文档混合检索⭐⭐⭐⭐⭐⭐高术语密集场景查询改写⭐⭐⭐⭐⭐⭐⭐高用户查询多变重排序⭐⭐⭐⭐⭐⭐⭐高高精度需求上下文压缩⭐⭐⭐⭐⭐中上下文长迭代检索⭐⭐⭐⭐⭐⭐⭐中复杂推理自我纠错⭐⭐⭐⭐⭐⭐低准确率优先查询路由⭐⭐⭐⭐⭐⭐⭐中多类型问题