1. RAG系统搭建全流程解析
在当今AI技术快速发展的背景下,让大语言模型具备专业领域的知识能力已成为刚需。RAG(检索增强生成)技术通过外挂知识库的方式,完美解决了大模型在专业领域"一本正经地胡说八道"的问题。本文将手把手带你从零搭建完整的RAG系统,掌握每个环节的核心技术与实现细节。
1.1 RAG技术核心价值
传统大语言模型存在三大痛点:知识更新滞后、专业领域知识不足、回答缺乏可追溯性。RAG技术通过以下方式完美解决这些问题:
- 实时知识更新:只需更新向量数据库,模型就能获取最新知识,无需重新训练
- 领域专业化:通过专属知识库让通用大模型变身领域专家
- 回答可验证:每个回答都能追溯到原始文档片段,确保可信度
- 成本可控:相比微调大模型,RAG的投入产出比更高
实际案例:某金融公司使用RAG系统后,客服回答准确率从68%提升至92%,平均响应时间缩短40%,同时大幅降低了人工复核的工作量。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 文档处理与向量化
2.1 多格式文档解析实战
不同格式的文档需要专门的解析方法。以下是经过生产验证的文档解析方案:
python复制import os
import PyPDF2
import docx
from typing import Union
class DocumentParser:
@staticmethod
def parse_pdf(file_path: str) -> str:
"""PDF解析优化版:处理扫描件、加密文件等特殊情况"""
text = []
try:
with open(file_path, 'rb') as f:
reader = PyPDF2.PdfReader(f)
if reader.is_encrypted:
reader.decrypt('') # 尝试空密码解密
for page in reader.pages:
page_text = page.extract_text() or ""
# 处理常见PDF格式问题
page_text = page_text.replace('\x0c', ' ') # 替换换页符
text.append(page_text.strip())
except Exception as e:
print(f"PDF解析失败:{file_path},错误:{str(e)}")
return '\n'.join(text)
@staticmethod
def parse_docx(file_path: str) -> str:
"""DOCX解析增强版:处理表格、页眉页脚等元素"""
try:
doc = docx.Document(file_path)
paragraphs = []
for para in doc.paragraphs:
if para.text.strip(): # 跳过空段落
paragraphs.append(para.text)
# 提取表格内容
for table in doc.tables:
for row in table.rows:
row_text = ' | '.join(cell.text for cell in row.cells)
paragraphs.append(row_text)
return '\n'.join(paragraphs)
except Exception as e:
print(f"DOCX解析失败:{file_path},错误:{str(e)}")
return ""
@staticmethod
def parse_txt(file_path: str) -> str:
"""文本文件解析:自动检测编码"""
encodings = ['utf-8', 'gbk', 'gb2312', 'iso-8859-1']
for encoding in encodings:
try:
with open(file_path, 'r', encoding=encoding) as f:
return f.read()
except UnicodeDecodeError:
continue
raise ValueError(f"无法解码文件:{file_path}")
关键改进点:
- PDF解析增加加密处理、特殊字符清理
- DOCX解析新增表格内容提取
- 文本文件支持多种编码自动检测
- 完善的异常处理机制
2.2 智能文本分块策略
文本分块是影响RAG效果的关键因素。经过大量实验,我们总结出以下分块策略:
python复制from typing import List
import re
class TextChunker:
def __init__(self, max_chunk_size=500, overlap=50):
self.max_chunk_size = max_chunk_size
self.overlap = overlap # 块间重叠字符数
def chunk_by_sentence(self, text: str) -> List[str]:
"""基于句子的智能分块"""
# 中文句子分割(考虑常见标点)
sentences = re.split(r'(?<=[。!?;;])\s+', text)
chunks = []
current_chunk = []
current_length = 0
for sent in sentences:
sent = sent.strip()
if not sent:
continue
sent_length = len(sent)
if current_length + sent_length > self.max_chunk_size and current_chunk:
chunks.append(' '.join(current_chunk))
# 保留重叠部分
overlap_start = max(0, len(current_chunk[-1]) - self.overlap)
current_chunk = [current_chunk[-1][overlap_start:]] if overlap_start > 0 else []
current_length = sum(len(s) for s in current_chunk)
current_chunk.append(sent)
current_length += sent_length
if current_chunk:
chunks.append(' '.join(current_chunk))
return chunks
def chunk_by_section(self, text: str) -> List[str]:
"""基于章节结构的语义分块"""
# 识别常见标题模式
section_pattern = r'(?:\n|^)(#+ .+?|\d+\.\d+ .+?|[一二三四五六七八九十]+、.+?)\n'
sections = re.split(section_pattern, text)
if len(sections) < 2:
return self.chunk_by_sentence(text)
chunks = []
# 第一个元素可能是空字符串
for i in range(1, len(sections), 2):
title = sections[i]
content = sections[i+1] if i+1 < len(sections) else ""
chunk = f"{title}\n{content}"
if len(chunk) > self.max_chunk_size * 1.5:
# 如果章节内容过长,再按句子分割
chunks.extend(self.chunk_by_sentence(chunk))
else:
chunks.append(chunk)
return chunks
分块策略选择建议:
- 技术文档:使用
chunk_by_section保留章节结构 - 叙述性内容:使用
chunk_by_sentence确保语义连贯 - 高度结构化数据:可考虑按表格/列表分块
经验之谈:分块大小应根据embedding模型调整。例如,使用OpenAI的text-embedding-ada-002时,建议块大小在300-500字符;使用all-MiniLM-L6-v2时,500-800字符效果更佳。
3. 向量数据库搭建与优化
3.1 ChromaDB深度配置
python复制import chromadb
from chromadb.config import Settings
from chromadb.utils import embedding_functions
class VectorDBManager:
def __init__(self, db_path: str = "chroma_db"):
# 生产环境推荐配置
self.client = chromadb.PersistentClient(
path=db_path,
settings=Settings(
chroma_db_impl="duckdb+parquet",
persist_directory=db_path,
anonymized_telemetry=False # 禁用数据收集
)
)
# 多embedding模型支持
self.embedding_models = {
"miniLM": embedding_functions.SentenceTransformerEmbeddingFunction(
model_name="all-MiniLM-L6-v2"
),
"bge-small": embedding_functions.SentenceTransformerEmbeddingFunction(
model_name="BAAI/bge-small-en-v1.5"
)
}
def create_collection(self, name: str, model_type: str = "miniLM"):
"""创建带优化的集合"""
assert model_type in self.embedding_models, f"不支持的模型类型:{model_type}"
return self.client.get_or_create_collection(
name=name,
embedding_function=self.embedding_models[model_type],
metadata={
"hnsw:space": "cosine", # 相似度计算方式
"hnsw:M": 16, # 构建参数M
"hnsw:ef_construction": 200 # 构建参数ef
}
)
性能优化参数说明:
hnsw:space:相似度计算方式,可选"cosine"(默认)、"l2"、"ip"hnsw:M:影响索引构建质量和内存使用(值越大精度越高但内存占用越大)hnsw:ef_construction:影响索引构建时的搜索范围
3.2 批量处理与元数据设计
python复制import hashlib
from datetime import datetime
from typing import List, Dict
class DocumentIndexer:
@staticmethod
def generate_doc_id(file_path: str) -> str:
"""生成稳定文档ID"""
file_hash = hashlib.md5(file_path.encode()).hexdigest()
return f"doc_{file_hash[:8]}"
@staticmethod
def generate_chunk_id(doc_id: str, chunk_num: int) -> str:
"""生成块ID"""
return f"{doc_id}_chunk_{chunk_num}"
def index_documents(self, collection, file_paths: List[str]):
"""批量索引文档"""
batch_size = 100
total_files = len(file_paths)
for i in range(0, total_files, batch_size):
batch_files = file_paths[i:i+batch_size]
print(f"处理文件 {i+1}-{min(i+batch_size, total_files)}/{total_files}")
all_ids = []
all_texts = []
all_metadatas = []
for file_path in batch_files:
try:
# 解析文档
text = DocumentParser().parse_document(file_path)
chunks = TextChunker().chunk_by_section(text)
# 生成元数据
doc_id = self.generate_doc_id(file_path)
file_name = os.path.basename(file_path)
file_size = os.path.getsize(file_path)
last_modified = datetime.fromtimestamp(
os.path.getmtime(file_path)
).isoformat()
# 准备批量插入数据
for j, chunk in enumerate(chunks):
chunk_id = self.generate_chunk_id(doc_id, j)
all_ids.append(chunk_id)
all_texts.append(chunk)
all_metadatas.append({
"doc_id": doc_id,
"chunk_num": j,
"file_name": file_name,
"file_size": file_size,
"last_modified": last_modified,
"content_type": self._detect_content_type(chunk)
})
except Exception as e:
print(f"文件处理失败:{file_path},错误:{str(e)}")
continue
# 批量插入
if all_ids:
collection.add(
documents=all_texts,
metadatas=all_metadatas,
ids=all_ids
)
print(f"已插入 {len(all_ids)} 个文本块")
@staticmethod
def _detect_content_type(text: str) -> str:
"""自动检测内容类型"""
if re.search(r'\d+\.\d+\.\d+', text): # 版本号
return "technical"
elif len(text.splitlines()) > 3: # 多行文本
return "structured"
elif len(text) < 100: # 短文本
return "summary"
return "general"
元数据设计最佳实践:
- 包含文档来源信息(文件路径、大小、修改时间)
- 记录块位置信息(文档ID、块序号)
- 添加内容特征标记(类型、关键词等)
- 保留业务相关属性(部门、项目、保密等级等)
4. 语义检索与结果优化
4.1 混合搜索策略
python复制class HybridSearcher:
def __init__(self, collection):
self.collection = collection
def semantic_search(self, query: str, top_k: int = 3,
filters: dict = None) -> dict:
"""带过滤条件的语义搜索"""
search_params = {
"query_texts": [query],
"n_results": top_k
}
if filters:
search_params["where"] = self._build_where_clause(filters)
return self.collection.query(**search_params)
def keyword_search(self, query: str, top_k: int = 3,
filters: dict = None) -> dict:
"""关键词搜索(使用元数据过滤)"""
# 提取关键词
keywords = self._extract_keywords(query)
if not keywords:
return {"documents": [], "metadatas": [], "distances": []}
# 构建where条件
where_conditions = []
for kw in keywords:
where_conditions.append({
"$or": [
{"file_name": {"$contains": kw}},
{"content_type": {"$contains": kw}}
]
})
if filters:
where_conditions.append(self._build_where_clause(filters))
return self.collection.query(
query_texts=[query],
n_results=top_k,
where={"$and": where_conditions} if where_conditions else None
)
def hybrid_search(self, query: str, top_k: int = 3,
filters: dict = None) -> dict:
"""混合搜索:结合语义和关键词"""
semantic_results = self.semantic_search(query, top_k, filters)
keyword_results = self.keyword_search(query, top_k, filters)
# 结果融合与去重
combined = {
"documents": semantic_results["documents"][0] + keyword_results["documents"][0],
"metadatas": semantic_results["metadatas"][0] + keyword_results["metadatas"][0],
"distances": semantic_results["distances"][0] + keyword_results["distances"][0]
}
# 按相似度排序
sorted_results = sorted(
zip(combined["documents"], combined["metadatas"], combined["distances"]),
key=lambda x: x[2]
)[:top_k]
# 解包排序后的结果
if sorted_results:
docs, metas, dists = zip(*sorted_results)
return {
"documents": [list(docs)],
"metadatas": [list(metas)],
"distances": [list(dists)]
}
return {"documents": [], "metadatas": [], "distances": []}
@staticmethod
def _build_where_clause(filters: dict) -> dict:
"""构建ChromaDB的where查询条件"""
where = {}
for k, v in filters.items():
if isinstance(v, list):
where[k] = {"$in": v}
else:
where[k] = {"$eq": v}
return where
@staticmethod
def _extract_keywords(query: str) -> List[str]:
"""简单关键词提取"""
# 移除标点
clean_query = re.sub(r'[^\w\s]', '', query)
# 提取名词短语(简单实现)
words = clean_query.split()
return [w for w in words if len(w) > 2]
搜索策略选择指南:
| 场景 | 推荐方法 | 优点 |
|---|---|---|
| 精确概念查询 | 关键词搜索 | 结果精准,速度快 |
| 语义相似查询 | 语义搜索 | 理解用户意图,灵活 |
| 综合查询 | 混合搜索 | 兼顾精准与灵活 |
| 带业务过滤 | 带where条件的搜索 | 结果更相关 |
4.2 结果重排序与增强
python复制class ResultEnhancer:
def __init__(self, llm_client):
self.llm_client = llm_client
def rerank_results(self, query: str, results: dict) -> dict:
"""使用LLM对搜索结果进行重排序"""
if not results or not results["documents"]:
return results
# 准备重排序prompt
passages = [
f"文档{i+1}:\n{text}\n相关度分数:{1-dist:.2f}"
for i, (text, dist) in enumerate(zip(
results["documents"][0],
results["distances"][0]
))
]
prompt = f"""根据用户问题和文档相关性,重新排序以下文档片段。
用户问题: {query}
文档片段:
{"\n\n".join(passages)}
请按相关性从高到低输出文档编号(如 2,1,3),不要解释:"""
try:
response = self.llm_client.chat.completions.create(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": prompt}],
temperature=0
)
new_order = [int(x)-1 for x in response.choices[0].message.content.split(",")]
# 按新顺序重组结果
return {
"documents": [[results["documents"][0][i] for i in new_order]],
"metadatas": [[results["metadatas"][0][i] for i in new_order]],
"distances": [[results["distances"][0][i] for i in new_order]]
}
except Exception as e:
print(f"重排序失败: {str(e)}")
return results
def add_citations(self, results: dict) -> str:
"""为结果添加引用来源"""
if not results or not results["documents"]:
return "未找到相关信息"
cited_text = []
for i, (doc, meta) in enumerate(zip(
results["documents"][0],
results["metadatas"][0]
)):
source = meta.get("file_name", "未知文档")
cited_text.append(f"[{i+1}] {doc}\n—— 来源: {source}")
return "\n\n".join(cited_text)
def highlight_keywords(self, text: str, keywords: List[str]) -> str:
"""在文本中高亮显示关键词"""
for kw in keywords:
text = text.replace(kw, f"**{kw}**")
return text
专业建议:在生产环境中,可以考虑使用专门的rerank模型(如cohere的rerank或bge的reranker),它们比通用LLM更擅长判断文档相关性。
5. LLM集成与对话管理
5.1 智能提示词工程
python复制class PromptEngineer:
@staticmethod
def build_rag_prompt(query: str, context: str,
chat_history: str = None) -> str:
"""构建RAG专用提示词"""
base_prompt = """你是一个专业的知识助手,请严格根据提供的上下文信息回答问题。
如果信息不足,请回答"根据现有信息无法确定",不要编造信息。"""
if chat_history:
base_prompt += f"\n\n之前的对话记录:\n{chat_history}"
prompt = f"""{base_prompt}
相关上下文:
{context}
当前问题: {query}
请用中文回答,保持专业但易懂。如果适用,请引用上下文中的具体数据。
回答:"""
return prompt
@staticmethod
def build_rewrite_prompt(query: str, history: str) -> str:
"""构建查询重写提示词"""
return f"""请将以下后续问题改写为完整的独立问题,参考之前的对话记录。
保持原意不变,但使其可以不依赖对话历史也能被理解。
对话历史:
{history}
后续问题: {query}
独立完整的问题:"""
@staticmethod
def build_summarize_prompt(text: str) -> str:
"""构建摘要提示词"""
return f"""请用1-2句话总结以下内容的核心信息,保留关键数据和结论。
内容:
{text}
摘要:"""
提示词设计原则:
- 明确角色:定义AI的专家身份
- 严格约束:限定回答必须基于给定上下文
- 对话连贯:融入历史对话记录
- 格式要求:指定回答语言和风格
- 安全边界:防止幻觉回答
5.2 对话状态管理
python复制import uuid
from datetime import datetime
from typing import Dict, List
class ConversationManager:
def __init__(self, max_history=10):
self.sessions: Dict[str, List[Dict]] = {}
self.max_history = max_history
def start_session(self, user_id: str = None) -> str:
"""开启新对话会话"""
session_id = user_id or str(uuid.uuid4())
self.sessions[session_id] = []
return session_id
def add_message(self, session_id: str, role: str, content: str):
"""记录对话消息"""
if session_id not in self.sessions:
self.start_session(session_id)
self.sessions[session_id].append({
"role": role,
"content": content,
"timestamp": datetime.now().isoformat()
})
# 保持对话历史不超过限制
if len(self.sessions[session_id]) > self.max_history:
self.sessions[session_id] = self.sessions[session_id][-self.max_history:]
def get_history(self, session_id: str,
last_n: int = None) -> List[Dict]:
"""获取对话历史"""
history = self.sessions.get(session_id, [])
return history[-last_n:] if last_n else history
def summarize_history(self, session_id: str, llm_client) -> str:
"""摘要长对话历史"""
history = self.get_history(session_id)
if not history:
return ""
conv_text = "\n".join(
f"{msg['role']}: {msg['content']}"
for msg in history
)
prompt = PromptEngineer.build_summarize_prompt(conv_text)
try:
response = llm_client.chat.completions.create(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": prompt}],
temperature=0.2
)
return response.choices[0].message.content
except Exception as e:
print(f"对话摘要失败: {str(e)}")
return ""
def contextualize_query(self, session_id: str,
query: str, llm_client) -> str:
"""上下文感知的查询重写"""
history = self.get_history(session_id, last_n=3)
if not history:
return query
conv_text = "\n".join(
f"{msg['role']}: {msg['content']}"
for msg in history
)
prompt = PromptEngineer.build_rewrite_prompt(query, conv_text)
try:
response = llm_client.chat.completions.create(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": prompt}],
temperature=0
)
return response.choices[0].message.content
except Exception as e:
print(f"查询重写失败: {str(e)}")
return query
对话管理高级技巧:
- 对长对话自动生成摘要,避免token超限
- 实现多轮指代消解(如"它"、"这个"等)
- 记录用户偏好(回答风格、详细程度等)
- 支持对话主题切换检测
- 添加敏感词过滤和内容审核
6. 系统集成与性能优化
6.1 完整RAG工作流实现
python复制class RAGSystem:
def __init__(self, vector_db, llm_client):
self.vector_db = vector_db
self.llm_client = llm_client
self.conv_manager = ConversationManager()
self.searcher = HybridSearcher(vector_db)
self.enhancer = ResultEnhancer(llm_client)
def query(self, question: str, session_id: str = None,
top_k: int = 3) -> dict:
"""端到端RAG查询流程"""
# 获取或创建会话
if not session_id:
session_id = self.conv_manager.start_session()
# 1. 上下文感知的查询重写
refined_query = self.conv_manager.contextualize_query(
session_id, question, self.llm_client
)
print(f"优化后的查询: {refined_query}")
# 2. 执行混合搜索
search_results = self.searcher.hybrid_search(refined_query, top_k)
# 3. 结果重排序与增强
reranked_results = self.enhancer.rerank_results(
refined_query, search_results
)
# 4. 构建提示词并调用LLM
context = self.enhancer.add_citations(reranked_results)
history = "\n".join(
f"{msg['role']}: {msg['content']}"
for msg in self.conv_manager.get_history(session_id)
)
prompt = PromptEngineer.build_rag_prompt(
question, context, history
)
try:
response = self.llm_client.chat.completions.create(
model="gpt-4-turbo",
messages=[{"role": "user", "content": prompt}],
temperature=0.3,
max_tokens=500
)
answer = response.choices[0].message.content
# 5. 记录对话历史
self.conv_manager.add_message(session_id, "user", question)
self.conv_manager.add_message(session_id, "assistant", answer)
return {
"answer": answer,
"sources": [
meta["file_name"]
for meta in reranked_results.get("metadatas", [[]])[0]
],
"session_id": session_id
}
except Exception as e:
print(f"LLM调用失败: {str(e)}")
return {
"answer": "系统暂时无法处理您的请求",
"sources": [],
"session_id": session_id
}
6.2 性能优化实战技巧
- 缓存机制:对常见查询结果进行缓存
python复制from functools import lru_cache
class CachedSearcher(HybridSearcher):
@lru_cache(maxsize=1000)
def semantic_search(self, query: str, top_k: int = 3) -> dict:
return super().semantic_search(query, top_k)
- 异步处理:使用异步IO提高吞吐量
python复制import asyncio
async def async_query(system: RAGSystem, questions: List[str]):
tasks = [
asyncio.create_task(
asyncio.to_thread(system.query, q)
) for q in questions
]
return await asyncio.gather(*tasks)
- 批量处理:文档索引批量操作
python复制def batch_index_files(system: RAGSystem, file_paths: List[str], batch_size=50):
for i in range(0, len(file_paths), batch_size):
batch = file_paths[i:i+batch_size]
with ThreadPoolExecutor() as executor:
executor.map(system.index_document, batch)
- 监控指标:关键性能指标跟踪
python复制class PerformanceMonitor:
metrics = {
"search_latency": [],
"llm_latency": [],
"cache_hit_rate": 0
}
@classmethod
def log_search_time(cls, duration: float):
cls.metrics["search_latency"].append(duration)
if len(cls.metrics["search_latency"]) > 100:
cls.metrics["search_latency"].pop(0)
@classmethod
def get_avg_search_time(cls) -> float:
if not cls.metrics["search_latency"]:
return 0
return sum(cls.metrics["search_latency"]) / len(cls.metrics["search_latency"])
7. 生产环境最佳实践
7.1 部署架构建议
中小规模部署方案:
code复制前端应用 → 负载均衡 → [RAG服务集群] → 向量数据库 → 大模型API
│
├─ 缓存层(Redis)
└─ 监控系统(Prometheus)
关键组件:
- 无状态服务:RAG服务应设计为无状态,方便横向扩展
- 读写分离:向量数据库的读写操作分离
- 分级缓存:查询结果缓存 + embedding缓存
- 熔断机制:当LLM API响应慢时自动降级
7.2 安全与权限控制
python复制class AccessController:
def __init__(self):
self.document_permissions = {} # 文档ID → 允许访问的角色
def check_permission(self, user_roles: List[str], doc_id: str) -> bool:
"""检查用户是否有权限访问文档"""
allowed_roles = self.document_permissions.get(doc_id, [])
return any(role in allowed_roles for role in user_roles)
def filter_results(self, results: dict, user_roles: List[str]) -> dict:
"""过滤无权限查看的结果"""
if not results or not results["metadatas"]:
return results
filtered_docs = []
filtered_metas = []
filtered_dists = []
for doc, meta, dist in zip(
results["documents"][0],
results["metadatas"][0],
results["distances"][0]
):
doc_id = meta.get("doc_id")
if not doc_id or self.check_permission(user_roles, doc_id):
filtered_docs.append(doc)
filtered_metas.append(meta)
filtered_dists.append(dist)
return {
"documents": [filtered_docs],
"metadatas": [filtered_metas],
"distances": [filtered_dists]
}
安全措施清单:
- 文档级访问控制
- 查询输入过滤(防注入攻击)
- 输出内容审核
- API调用限流
- 敏感数据脱敏
7.3 持续维护策略
-
知识库更新机制:
- 定期增量更新
- 文档变更监听
- 版本控制集成
-
效果监控体系:
python复制class QualityMonitor: def log_feedback(self, query: str, response: str, useful: bool, user_comment: str = ""): """记录用户反馈""" # 存储到数据库或日志系统 pass def calculate_accuracy(self) -> float: """计算回答准确率""" # 基于用户反馈计算 pass -
迭代优化流程:
code复制收集反馈 → 分析问题 → 优化组件 → A/B测试 → 全量发布 ↑ ↓ └────── 监控效果 ←──────┘
8. 常见问题与解决方案
8.1 典型错误排查表
| 问题现象 | 可能原因 | 解决方案 |
|---|---|---|
| 返回无关内容 | 1. 分块大小不合适 2. embedding模型不匹配 |
1. 调整分块策略 2. 更换embedding模型 |
| 回答不准确 | 1. 检索结果差 2. prompt设计问题 |
1. 优化检索策略 2. 改进prompt约束 |
| 响应速度慢 | 1. 网络延迟 2. 未使用缓存 |
1. 检查网络 2. 添加缓存层 |
| 内存占用高 | 1. 向量索引配置不当 2. 批量太大 |
1. 调整HNSW参数 2. 减小批量大小 |
8.2 高级调试技巧
-
检索可视化:使用PCA/t-SNE降维展示查询与文档的向量关系
python复制from sklearn.decomposition import PCA import matplotlib.pyplot as plt def plot_embeddings(query_embedding, doc_embeddings): """可视化embedding空间""" all_embeddings = [query_embedding] + doc_embeddings pca = PCA(n_components=2) points = pca.fit_transform(all_embeddings) plt.figure(figsize=(10, 6)) plt.scatter(points[1:, 0], points[1:, 1], label='文档') plt.scatter(points[0, 0], points[0, 1], c='red', label='查询') plt.legend() plt.show() -
搜索过程记录:
python复制class SearchDebugger: def __init__(self): self.search_log = [] def log_search(self, query: str, results: dict): """记录搜索详情""" self.search_log.append({ "query": query, "top_match": results["documents"][0][0] if results["documents"] else None, "score": 1 - results["distances"][0][0] if results["distances"] else 0 }) -
AB测试框架:
python复制class ABTester: def compare_strategies(self, queries: List[str], strategy_a: callable, strategy_b: callable): """对比两种搜索策略""" results = [] for q in queries: res_a = strategy_a(q) res_b = strategy_b(q) results.append({ "query": q, "a_score": self._calculate_relevance(q, res_a), "b_score": self._calculate_relevance(q, res_b) }) return results
9. 扩展应用与进阶方向
9.1 多模态RAG实现
python复制class MultiModalRAG:
def __init__(self, text_db, image_db):
self.text_retriever = HybridSearcher(text_db)
self.image_retriever = ImageSearcher(image_db)
def search(self, query: str, image: bytes = None):
"""多模态搜索"""
text_results = self.text_retriever.semantic_search(query)
if image:
image_results = self.image_retriever.similarity_search(image)
return self._fuse_results(text_results, image_results)
return text_results
@staticmethod
def _fuse_results(text_res: dict, image_res: dict) -> dict:
"""融合文本和图像搜索结果"""
# 实现跨模态结果融合算法
pass
9.2 实时知识更新方案
python复制import watchdog.events
import watchdog.observers
class KnowledgeUpdater(watchdog.events.FileSystemEventHandler):
def __init__(self, rag_system: RAGSystem):
self.rag = rag_system
self.observer = watchdog.observers.Observer()
def on_modified(self, event):
"""监听文件变更"""
if not event.is_directory:
print(f"检测到文件变更: {event.src_path
