AI模块化RAG:单体不够用了之后
结果遇到了一堆问题:
- 检索不准:有时候用户问个具体问题,回来的都是些泛泛而谈的内容
- 上下文过长:为了提高召回率,把很多文档都塞进去了,结果 token 不够用
- 无法优化:每次想改个检索策略,都得把整个流程重构一遍
后来发现,原来我们用的这种"一条龙"的 RAG 模式叫单体 RAG(Monolithic RAG)。
踩坑背景
最近在做一个企业的知识库项目,本来以为把文档丢进向量数据库,然后加上个简单的检索+生成就能完事。结果遇到了一堆问题:
- 检索不准:有时候用户问个具体问题,回来的都是些泛泛而谈的内容
- 上下文过长:为了提高召回率,把很多文档都塞进去了,结果 token 不够用
- 无法优化:每次想改个检索策略,都得把整个流程重构一遍
后来发现,原来我们用的这种"一条龙"的 RAG 模式叫单体 RAG(Monolithic RAG)。这种方式简单是简单,但扩展性和可维护性都不太行。
于是我就开始研究模块化 RAG(Modular RAG),想把整个流程拆解开来,让它更灵活、更好维护。这篇文章就是我在实践模块化 RAG 过程中的踩坑和总结。
为什么要模块化
单体 RAG 的痛点
我先说说单体 RAG 的问题。典型的单体 RAG 长这样:
这种模式在简单场景下够用,但复杂问题就暴露了:
- 检索方式单一:只能用向量检索,对一些精确匹配的问题效果很差
- 无法根据问题调整策略:不管用户问什么,都用同一套流程
- 难以 A/B 测试:想测试不同的检索策略,得复制整个系统
模块化 RAG 的优势
模块化 RAG 的核心思想是:把 RAG 流程拆分成独立的模块,每个模块都可以单独替换和优化。
这样的好处:
- 每个模块可以独立开发和测试
- 可以根据问题类型选择不同的处理流程
- 容易扩展新的功能模块
实现方案
模块设计
我设计的模块化 RAG 包含以下几个核心模块:
1. 问题路由模块(Query Router)
这个模块的作用是分析用户问题,然后决定用哪种检索策略。
from typing import Literal
from pydantic import BaseModel
class QueryType(BaseModel):
query_type: Literal["factual", "analytical", "creative"]
confidence: float
def route_query(question: str) -> QueryType:
"""
根据问题类型决定路由
"""
prompt = f"""
分析以下问题的类型:
问题:{question}
类型说明:
- factual: 事实性查询,需要精确信息
- analytical: 分析性查询,需要综合多个信息
- creative: 创造性查询,需要发散思维
返回类型和置信度(0-1)。
"""
# 这里调用 LLM 进行分类
result = llm.predict(prompt, response_model=QueryType)
return result
2. 检索模块(Retrieval)
检索模块支持多种检索方式:
from abc import ABC, abstractmethod
class BaseRetriever(ABC):
@abstractmethod
def retrieve(self, query: str, top_k: int = 5) -> list[dict]:
pass
class VectorRetriever(BaseRetriever):
def __init__(self, vector_db):
self.vector_db = vector_db
def retrieve(self, query: str, top_k: int = 5) -> list[dict]:
# 向量检索实现
embeddings = self.embed(query)
results = self.vector_db.similarity_search(
embeddings, k=top_k * 2 # 多召回一些用于重排序
)
return results
class KeywordRetriever(BaseRetriever):
def __init__(self, index):
self.index = index
def retrieve(self, query: str, top_k: int = 5) -> list[dict]:
# 关键词检索实现
keywords = self.extract_keywords(query)
results = self.index.search(keywords, top_k=top_k * 2)
return results
class HybridRetriever(BaseRetriever):
def __init__(self, retrievers: list[BaseRetriever], weights: list[float]):
self.retrievers = retrievers
self.weights = weights
def retrieve(self, query: str, top_k: int = 5) -> list[dict]:
# 混合检索
all_results = []
for retriever, weight in zip(self.retrievers, self.weights):
results = retriever.retrieve(query, top_k=top_k)
for result in results:
result['score'] *= weight
all_results.extend(results)
# 合并和去重
merged = self.merge_results(all_results)
return merged[:top_k * 2]
3. 重排序模块(Reranker)
这个模块对检索到的结果进行重新排序,提高相关性。
class Reranker:
def __init__(self, model_name: str = "BAAI/bge-reranker-base"):
from sentence_transformers import CrossEncoder
self.model = CrossEncoder(model_name)
def rerank(self, query: str, documents: list[dict], top_k: int = 5) -> list[dict]:
"""
对文档进行重排序
"""
# 准备输入对
pairs = [(query, doc['content']) for doc in documents]
# 计算相关性分数
scores = self.model.predict(pairs)
# 更新文档分数
for doc, score in zip(documents, scores):
doc['rerank_score'] = score
# 按新分数排序
documents.sort(key=lambda x: x['rerank_score'], reverse=True)
return documents[:top_k]
4. 上下文组装模块(Context Builder)
根据重排序后的结果组装上下文。
class ContextBuilder:
def __init__(self, max_tokens: int = 3000):
self.max_tokens = max_tokens
def build(self, query: str, documents: list[dict]) -> str:
"""
组装上下文,控制 token 长度
"""
context_parts = []
current_tokens = 0
for doc in documents:
# 估算 token 数(简单实现)
doc_tokens = len(doc['content']) // 4 # 粗略估算
if current_tokens + doc_tokens > self.max_tokens:
break
context_parts.append(f"""
## 来源:{doc.get('source', '未知')}
{doc['content']}
""")
current_tokens += doc_tokens
return "\n".join(context_parts)
5. 生成模块(Generator)
最后用 LLM 生成答案。
class Generator:
def __init__(self, model_name: str = "gpt-4"):
self.model_name = model_name
def generate(self, query: str, context: str) -> str:
"""
根据问题和上下文生成答案
"""
prompt = f"""
基于以下上下文回答问题:
上下文:
{context}
问题:{query}
要求:
1. 只基于上下文回答,不要编造信息
2. 如果上下文没有相关信息,明确说明
3. 回答要准确、简洁
"""
response = llm.predict(prompt, model=self.model_name)
return response
整合所有模块
现在把这些模块整合起来:
class ModularRAG:
def __init__(self, config: dict):
self.query_router = QueryRouter()
self.retrievers = self._init_retrievers(config)
self.reranker = Reranker()
self.context_builder = ContextBuilder(config.get('max_tokens', 3000))
self.generator = Generator(config.get('generator_model', 'gpt-4'))
def _init_retrievers(self, config: dict):
retrievers = {}
if config.get('vector_retriever'):
retrievers['vector'] = VectorRetriever(config['vector_retriever'])
if config.get('keyword_retriever'):
retrievers['keyword'] = KeywordRetriever(config['keyword_retriever'])
return retrievers
def query(self, question: str) -> dict:
"""
完整的查询流程
"""
# 1. 路由问题
query_type = self.query_router.route(question)
# 2. 选择检索器
retriever = self._select_retriever(query_type)
# 3. 检索
documents = retriever.retrieve(question, top_k=10)
# 4. 重排序
documents = self.reranker.rerank(question, documents, top_k=5)
# 5. 组装上下文
context = self.context_builder.build(question, documents)
# 6. 生成答案
answer = self.generator.generate(question, context)
return {
'answer': answer,
'sources': documents,
'query_type': query_type
}
def _select_retriever(self, query_type: QueryType) -> BaseRetriever:
"""
根据问题类型选择检索器
"""
if query_type.query_type == 'factual':
return HybridRetriever([
self.retrievers['vector'],
self.retrievers['keyword']
], [0.7, 0.3])
elif query_type.query_type == 'analytical':
return self.retrievers['vector']
else:
return self.retrievers['vector']
踩坑记录
坑 1:路由准确性不够
刚开始用简单的关键词匹配来路由问题,结果经常分错。比如用户问"我们的产品有什么特色?",这不是 factual 问题,但系统把它当成 factual 了。
解决方案:改用 LLM 进行分类,并增加置信度检查。
def route_query_with_confidence(question: str) -> tuple[QueryType, bool]:
"""
带置信度检查的路由
"""
result = route_query(question)
# 如果置信度太低,使用默认策略
if result.confidence < 0.7:
result.query_type = "factual" # 默认用更稳妥的策略
return result, result.confidence >= 0.7
坑 2:检索结果过多
一开始为了提高召回率,检索了很多文档(比如 top_k=20),结果导致:
- 重排序很慢
- 上下文太长,token 不够用
解决方案:
- 控制检索数量(top_k=10)
- 在重排序后再筛选(top_k=5)
坑 3:重排序引入延迟
用了 BGE-Reranker 之后,虽然准确率提高了,但延迟增加了 200-300ms。
解决方案:
- 缓存重排序结果
- 对于高频问题,可以提前预计算
from functools import lru_cache
class CachedReranker(Reranker):
@lru_cache(maxsize=1000)
def rerank_cached(self, query: str, docs_tuple: tuple, top_k: int = 5) -> list[dict]:
"""
缓存的重排序
"""
docs = list(docs_tuple)
return super().rerank(query, docs, top_k)
def rerank(self, query: str, documents: list[dict], top_k: int = 5) -> list[dict]:
# 转为 tuple 以便缓存
docs_tuple = tuple(doc['content'] for doc in documents)
return self.rerank_cached(query, docs_tuple, top_k)
坑 4:上下文截断导致信息丢失
组装上下文时,如果直接截断,可能会把重要的信息截掉。
解决方案:
- 按相关性排序后再截断
- 对于重要文档,保留完整内容
class SmartContextBuilder(ContextBuilder):
def build(self, query: str, documents: list[dict]) -> str:
"""
智能组装上下文
"""
# 按相关性排序
documents = sorted(
documents,
key=lambda x: x.get('rerank_score', 0),
reverse=True
)
# 前 3 个文档保留完整
important_docs = documents[:3]
remaining_docs = documents[3:]
context_parts = []
current_tokens = 0
# 先添加重要文档
for doc in important_docs:
doc_content = f"## 来源:{doc.get('source', '未知')}\n{doc['content']}\n"
doc_tokens = len(doc_content) // 4
if current_tokens + doc_tokens > self.max_tokens * 0.8:
break # 为剩余文档留空间
context_parts.append(doc_content)
current_tokens += doc_tokens
# 添加剩余文档(摘要)
for doc in remaining_docs:
if current_tokens >= self.max_tokens:
break
summary = self._summarize(doc['content'])
context_parts.append(f"## 来源:{doc.get('source', '未知')}\n{summary}\n")
current_tokens += len(summary) // 4
return "\n".join(context_parts)
实践效果
实施模块化 RAG 后,效果对比:
| 指标 | 单体 RAG | 模块化 RAG | 提升 |
|---|---|---|---|
| 准确率 | 72% | 85% | +13% |
| 平均延迟 | 1.2s | 1.5s | +300ms |
| Token 使用 | 2500 | 1800 | -28% |
| 开发效率 | 基线 | +40% | 更快迭代 |
具体改进
- 准确率提升:通过混合检索和重排序,相关性明显提高
- Token 优化:智能组装上下文,避免浪费
- 可维护性:每个模块独立,方便调试和优化
- 扩展性:可以轻松添加新的检索器或处理器
使用示例
# 配置
config = {
'vector_retriever': 'chromadb',
'keyword_retriever': 'elasticsearch',
'max_tokens': 3000,
'generator_model': 'gpt-4'
}
# 初始化
rag = ModularRAG(config)
# 查询
result = rag.query("我们的产品有哪些特色?")
print(f"答案:{result['answer']}")
print(f"问题类型:{result['query_type']}")
print(f"参考来源:{len(result['sources'])} 个文档")
后续优化方向
- 查询改写:对用户问题进行改写,提高检索准确率
- 多轮对话:支持上下文记忆和追问
- 反馈学习:根据用户反馈调整检索策略
- 性能优化:并行处理、缓存优化
结语
模块化 RAG 让我们从"一条龙"的简单模式,转向了更灵活、可维护的架构。虽然初期投入会多一些,但从长期来看,带来的收益是明显的:
- 开发效率提高,迭代更快
- 问题定位更容易,调试更简单
- 功能扩展更灵活,不用推倒重来
如果你也在做 RAG 相关的项目,建议一开始就考虑模块化设计,避免后期重构的痛苦。当然,具体怎么模块化,还是要根据你的实际需求来决定,不要过度设计。
希望这篇文章能帮到同样在 RAG 道路上摸索的同学,有问题欢迎交流!
版权声明: 本文首发于 指尖魔法屋-AI模块化RAG:单体不够用了之后(https://blog.thinkmoon.cn/post/379-ai-modular-rag-monolithic-modular-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!
评论
使用 GitHub 账号登录后即可留言,支持 Markdown。