目录
- RAG 架构全景
- Embedding 模型选型
- 文档分块策略
- 向量检索与混合搜索
- 重排序(Reranker)
- LangChain RAG 实现
- LlamaIndex RAG 实现
- 端到端 FastAPI RAG 服务
- RAG 评估与优化
1. RAG 架构全景
RAG(Retrieval-Augmented Generation)通过将外部知识检索与 LLM 生成结合,解决 hallucination(幻觉)问题。
1.1 标准 RAG 流程
用户查询 → Query Embedding → 向量检索 → Top-K 文档 → Rerank → 上下文拼接 → LLM 生成
from dataclasses import dataclass
from typing import List
@dataclass
class RAGContext:
query: str
retrieved_chunks: List[str]
scores: List[float]
context_window: str # 拼接后的上下文
class RAGPipeline:
"""标准 RAG 流水线。"""
def __init__(self, embedder, retriever, llm_client):
self.embedder = embedder
self.retriever = retriever
self.llm = llm_client
async def answer(self, query: str, top_k: int = 5) -> str:
# 1. Query → Embedding
query_embedding = self.embedder.encode(query)
# 2. 检索 Top-K
chunks, scores = self.retriever.search(query_embedding, top_k)
# 3. 构建上下文
context = "\n\n---\n\n".join(
f"[Document {i+1}]\n{chunk}" for i, chunk in enumerate(chunks)
)
# 4. LLM 生成
prompt = f"""Answer the question based on the provided context.
If the answer is not in the context, say "I don't know".
Context:
{context}
Question: {query}
Answer:"""
return await self.llm.chat_completion([{"role": "user", "content": prompt}])
2. Embedding 模型选型
| 模型 | 维度 | 最大长度 | 语言 | 性能(MTEB) | 许可 | 适用 |
|---|---|---|---|---|---|---|
| text-embedding-3-small | 1536 | 8192 | 多语言 | 62.3% | 商业 | 通用,低成本 |
| text-embedding-3-large | 3072 | 8192 | 多语言 | 64.6% | 商业 | 高质量需求 |
| BGE-M3 | 1024 | 8192 | 多语言 | 72.3% | MIT | 开源首选 |
| GTE-large | 1024 | 512 | 多语言 | 69.3% | MIT | 平衡速度与效果 |
| E5-mistral-7b-instruct | 4096 | 32768 | 多语言 | 74.1% | MIT | 长文档,高质量 |
| multilingual-e5-large | 1024 | 512 | 多语言 | 68.3% | MIT | 中文场景 |
2.1 本地 Embedding 推理
from sentence_transformers import SentenceTransformer
import numpy as np
class LocalEmbedder:
"""基于 sentence-transformers 的本地嵌入模型。"""
def __init__(self, model_name: str = "BAAI/bge-m3"):
self.model = SentenceTransformer(model_name)
self.dim = self.model.get_sentence_embedding_dimension()
def encode(self, texts: list[str] | str, normalize: bool = True) -> np.ndarray:
if isinstance(texts, str):
texts = [texts]
# BGE 模型需要添加指令前缀
instruction = "Represent this sentence for searching relevant passages: "
texts = [instruction + t for t in texts]
embeddings = self.model.encode(texts, convert_to_numpy=True)
if normalize:
embeddings = embeddings / np.linalg.norm(embeddings, axis=1, keepdims=True)
return embeddings
def encode_queries(self, queries: list[str]) -> np.ndarray:
"""Query 使用不同指令。"""
instruction = "Represent this query for retrieving relevant documents: "
texts = [instruction + q for q in queries]
embeddings = self.model.encode(texts, convert_to_numpy=True)
return embeddings / np.linalg.norm(embeddings, axis=1, keepdims=True)
2.2 多模态 Embedding(图像+文本)
from sentence_transformers import SentenceTransformer
class MultimodalEmbedder:
"""CLIP 风格的多模态嵌入。"""
def __init__(self):
self.model = SentenceTransformer("clip-ViT-B-32")
def encode_image(self, image_path: str) -> np.ndarray:
from PIL import Image
img = Image.open(image_path)
return self.model.encode(img)
def encode_text(self, text: str) -> np.ndarray:
return self.model.encode(text)
3. 文档分块策略
分块质量直接影响 RAG 效果——太大丢失细节,太小丢失上下文。
| 策略 | 方法 | 适用 | 缺点 |
|---|---|---|---|
| 固定大小 | 每 N 个 token/字符 | 简单场景 | 句子截断 |
| 递归文本 | 按段落→句子→词递归切分 | 通用文档 | 实现复杂 |
| 语义分块 | 按语义边界(句号/主题变化) | 高质量需求 | 计算开销 |
| 重叠窗口 | 固定大小 + M token 重叠 | 避免边界信息丢失 | 冗余增加 |
| 代理分块 | LLM 判断最佳切分点 | 最高质量 | Token 成本高 |
3.1 递归文本分块(推荐)
from typing import List
import re
class RecursiveTextSplitter:
"""LangChain 风格递归文本分块。"""
def __init__(
self,
chunk_size: int = 500,
chunk_overlap: int = 50,
separators: List[str] = None,
):
self.chunk_size = chunk_size
self.chunk_overlap = chunk_overlap
self.separators = separators or ["\n\n", "\n", ". ", "? ", "! ", " ", ""]
def split_text(self, text: str) -> List[str]:
return self._split_recursive(text, self.separators)
def _split_recursive(self, text: str, separators: List[str]) -> List[str]:
if len(text) <= self.chunk_size:
return [text] if text.strip() else []
separator = separators[0] if separators else ""
if separator:
parts = text.split(separator)
else:
parts = list(text)
chunks = []
current = ""
for part in parts:
candidate = current + (separator if current else "") + part
if len(candidate) <= self.chunk_size:
current = candidate
else:
if current:
chunks.append(current)
# 保留重叠部分
if self.chunk_overlap > 0:
current = current[-self.chunk_overlap:] + separator + part
else:
current = part
else:
# 单个部分超过 chunk_size,递归拆分
sub = self._split_recursive(part, separators[1:])
chunks.extend(sub)
current = ""
if current:
chunks.append(current)
return chunks
# 按 Markdown 标题智能分块
class MarkdownSplitter:
def __init__(self, chunk_size: int = 1000):
self.chunk_size = chunk_size
def split(self, markdown: str) -> List[dict]:
"""按标题层级分块,保留标题元数据。"""
chunks = []
current_lines = []
current_headers = []
for line in markdown.split("\n"):
if line.startswith("#"):
if current_lines:
chunks.append({
"content": "\n".join(current_lines),
"headers": current_headers.copy(),
})
level = len(line.split()[0])
current_headers = current_headers[:level-1] + [line.strip()]
current_lines = []
else:
current_lines.append(line)
if current_lines:
chunks.append({
"content": "\n".join(current_lines),
"headers": current_headers,
})
return chunks
3.2 按 Token 计数精确分块
import tiktoken
class TokenAwareSplitter:
def __init__(self, chunk_size: int = 512, chunk_overlap: int = 50, model: str = "gpt-4o"):
self.chunk_size = chunk_size
self.chunk_overlap = chunk_overlap
self.encoding = tiktoken.encoding_for_model(model)
def split(self, text: str) -> List[str]:
tokens = self.encoding.encode(text)
chunks = []
start = 0
while start < len(tokens):
end = min(start + self.chunk_size, len(tokens))
chunk_tokens = tokens[start:end]
chunks.append(self.encoding.decode(chunk_tokens))
start += self.chunk_size - self.chunk_overlap
return chunks
4. 向量检索与混合搜索
4.1 Dense 向量检索(Qdrant)
from qdrant_client import QdrantClient
from qdrant_client.models import Distance, VectorParams, PointStruct
class QdrantStore:
"""Qdrant 向量存储封装。"""
def __init__(self, host: str = "localhost", port: int = 6333, collection: str = "documents"):
self.client = QdrantClient(host=host, port=port)
self.collection = collection
def create_collection(self, dim: int):
self.client.recreate_collection(
collection_name=self.collection,
vectors_config=VectorParams(size=dim, distance=Distance.COSINE),
)
def upsert(self, chunks: list[str], embeddings: np.ndarray, metadata: list[dict] = None):
points = [
PointStruct(
id=i,
vector=emb.tolist(),
payload={"text": chunk, **(metadata[i] if metadata else {})},
)
for i, (emb, chunk) in enumerate(zip(embeddings, chunks))
]
self.client.upsert(collection_name=self.collection, points=points)
def search(self, query_embedding: np.ndarray, top_k: int = 5) -> list[dict]:
results = self.client.search(
collection_name=self.collection,
query_vector=query_embedding.tolist(),
limit=top_k,
)
return [
{"text": r.payload["text"], "score": r.score, "metadata": r.payload}
for r in results
]
4.2 Hybrid Search:Dense + Sparse(BM25)
from qdrant_client.models import SparseVector
class HybridRetriever:
"""混合检索:Dense 语义 + Sparse 词汇匹配。"""
def __init__(self, qdrant_store, sparse_model):
self.dense_store = qdrant_store
self.sparse_model = sparse_model # 如 splade或 bm25
def search(self, query: str, query_emb: np.ndarray, top_k: int = 10) -> list[dict]:
# Dense 检索
dense_results = self.dense_store.search(query_emb, top_k=top_k)
# Sparse 检索(BM25)
sparse_results = self.sparse_model.search(query, top_k=top_k)
# Reciprocal Rank Fusion (RRF) 融合
return self._rrf_fusion(dense_results, sparse_results, k=60)
def _rrf_fusion(self, dense: list, sparse: list, k: int = 60) -> list[dict]:
"""RRF:基于排名的融合算法。"""
scores = {}
for rank, r in enumerate(dense):
doc_id = r["metadata"].get("doc_id", r["text"][:50])
scores[doc_id] = scores.get(doc_id, 0) + 1 / (k + rank + 1)
if doc_id not in scores:
scores[doc_id] = {"text": r["text"], "metadata": r["metadata"]}
for rank, r in enumerate(sparse):
doc_id = r["metadata"].get("doc_id", r["text"][:50])
if doc_id in scores:
if isinstance(scores[doc_id], dict) and "text" in scores[doc_id]:
scores[doc_id] = 1 / (k + rank + 1)
else:
scores[doc_id] += 1 / (k + rank + 1)
else:
scores[doc_id] = 1 / (k + rank + 1)
scores[doc_id] = {"text": r["text"], "metadata": r["metadata"], "score": scores[doc_id]}
# 排序返回
sorted_docs = sorted(
[(k, v) for k, v in scores.items() if isinstance(v, (int, float)) or "score" in v],
key=lambda x: x[1] if isinstance(x[1], (int, float)) else x[1]["score"],
reverse=True,
)
return [doc for _, doc in sorted_docs[:10]]
5. 重排序(Reranker)
检索的 Top-K 结果顺序可能不够精确,使用轻量级 Reranker 模型对 Top-K 重新排序。
from sentence_transformers import CrossEncoder
class Reranker:
"""交叉编码器重排序。"""
def __init__(self, model_name: str = "BAAI/bge-reranker-large"):
self.model = CrossEncoder(model_name)
def rerank(self, query: str, documents: list[str], top_k: int = 5) -> list[tuple[str, float]]:
"""对文档列表重排序,返回 (doc, score)。"""
pairs = [(query, doc) for doc in documents]
scores = self.model.predict(pairs)
ranked = sorted(
zip(documents, scores),
key=lambda x: x[1],
reverse=True,
)
return ranked[:top_k]
典型 RAG 检索+重排序流程:
async def advanced_retrieve(query: str, top_k: int = 5) -> list[str]:
embedder = LocalEmbedder()
retriever = QdrantStore()
reranker = Reranker()
# 1. Query embedding
query_emb = embedder.encode_queries([query])
# 2. 粗排:向量检索 Top-20
candidates = retriever.search(query_emb[0], top_k=20)
candidate_texts = [c["text"] for c in candidates]
# 3. 精排:Reranker Top-5
ranked = reranker.rerank(query, candidate_texts, top_k=top_k)
return [doc for doc, _ in ranked]
6. LangChain RAG 实现
from langchain_community.vectorstores import Qdrant
from langchain_community.embeddings import HuggingFaceEmbeddings
from langchain_community.chat_models import ChatOpenAI
from langchain.chains import RetrievalQA
from langchain.prompts import PromptTemplate
class LangChainRAG:
def __init__(self, qdrant_url: str):
embeddings = HuggingFaceEmbeddings(model_name="BAAI/bge-m3")
self.vectorstore = Qdrant.from_existing_collection(
embedding=embeddings,
collection_name="documents",
url=qdrant_url,
)
llm = ChatOpenAI(model="gpt-4o", temperature=0)
prompt = PromptTemplate.from_template("""
You are a helpful assistant. Use the following context to answer the question.
If you don't know, say "I don't know".
Context:
{context}
Question: {question}
Answer:""")
self.qa_chain = RetrievalQA.from_chain_type(
llm=llm,
chain_type="stuff",
retriever=self.vectorstore.as_retriever(search_kwargs={"k": 5}),
chain_type_kwargs={"prompt": prompt},
return_source_documents=True,
)
async def ask(self, question: str) -> dict:
result = await self.qa_chain.ainvoke({"query": question})
return {
"answer": result["result"],
"sources": [d.page_content for d in result["source_documents"]],
}
7. LlamaIndex RAG 实现
from llama_index.core import VectorStoreIndex, SimpleDirectoryReader, Settings
from llama_index.embeddings.huggingface import HuggingFaceEmbedding
from llama_index.llms.openai import OpenAI
from llama_index.core.postprocessor import SentenceTransformerRerank
class LlamaIndexRAG:
def __init__(self, data_dir: str):
Settings.embed_model = HuggingFaceEmbedding(model_name="BAAI/bge-m3")
Settings.llm = OpenAI(model="gpt-4o")
documents = SimpleDirectoryReader(data_dir).load_data()
self.index = VectorStoreIndex.from_documents(documents)
# 添加 Reranker
self.rerank = SentenceTransformerRerank(
model="BAAI/bge-reranker-large",
top_n=5,
)
def query(self, question: str) -> str:
query_engine = self.index.as_query_engine(
similarity_top_k=20,
node_postprocessors=[self.rerank],
)
response = query_engine.query(question)
return str(response)
8. 端到端 FastAPI RAG 服务
from fastapi import FastAPI, UploadFile, File, HTTPException
from pydantic import BaseModel
import tempfile
import shutil
from pathlib import Path
app = FastAPI(title="RAG Knowledge Base API")
class QueryRequest(BaseModel):
question: str
top_k: int = 5
rerank: bool = True
class QueryResponse(BaseModel):
answer: str
sources: list[dict]
latency_ms: float
class IngestRequest(BaseModel):
text: str
metadata: dict = {}
class DocumentIngestor:
def __init__(self):
self.embedder = LocalEmbedder()
self.splitter = RecursiveTextSplitter(chunk_size=500, chunk_overlap=50)
self.store = QdrantStore()
self.store.create_collection(self.embedder.dim)
self.reranker = Reranker()
async def ingest_text(self, text: str, metadata: dict = None):
chunks = self.splitter.split_text(text)
embeddings = self.embedder.encode(chunks)
metas = [metadata or {} for _ in chunks]
self.store.upsert(chunks, embeddings, metas)
return len(chunks)
async def ingest_file(self, file: UploadFile):
"""支持 txt, md, pdf 文件上传。"""
suffix = Path(file.filename).suffix
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
shutil.copyfileobj(file.file, tmp)
tmp_path = tmp.name
try:
if suffix == ".pdf":
text = self._extract_pdf(tmp_path)
else:
with open(tmp_path, "r", encoding="utf-8") as f:
text = f.read()
return await self.ingest_text(text, {"source": file.filename})
finally:
Path(tmp_path).unlink()
def _extract_pdf(self, path: str) -> str:
from pypdf import PdfReader
reader = PdfReader(path)
return "\n".join(page.extract_text() or "" for page in reader.pages)
ingestor = DocumentIngestor()
@app.post("/ingest/text")
async def ingest_text(req: IngestRequest):
count = await ingestor.ingest_text(req.text, req.metadata)
return {"chunks_ingested": count}
@app.post("/ingest/file")
async def ingest_file(file: UploadFile = File(...)):
count = await ingestor.ingest_file(file)
return {"chunks_ingested": count, "filename": file.filename}
@app.post("/query", response_model=QueryResponse)
async def query(req: QueryRequest):
import time
start = time.perf_counter()
embedder = LocalEmbedder()
store = QdrantStore()
# 检索
query_emb = embedder.encode_queries([req.question])
results = store.search(query_emb[0], top_k=20 if req.rerank else req.top_k)
texts = [r["text"] for r in results]
sources = [{"text": r["text"], "score": r["score"]} for r in results]
# 重排序
if req.rerank:
reranked = ingestor.reranker.rerank(req.question, texts, top_k=req.top_k)
texts = [doc for doc, _ in reranked]
sources = [{"text": doc, "score": float(score)} for doc, score in reranked]
# 生成
context = "\n\n---\n\n".join(f"[{i+1}] {t}" for i, t in enumerate(texts))
from openai import AsyncOpenAI
oai = AsyncOpenAI()
response = await oai.chat.completions.create(
model="gpt-4o",
messages=[
{"role": "system", "content": "Answer based on the provided context only."},
{"role": "user", "content": f"Context:\n{context}\n\nQuestion: {req.question}"},
],
temperature=0.0,
)
latency = (time.perf_counter() - start) * 1000
return QueryResponse(
answer=response.choices[0].message.content,
sources=sources,
latency_ms=round(latency, 2),
)
9. RAG 评估与优化
9.1 评估指标
| 指标 | 含义 | 计算方法 |
|---|---|---|
| Context Precision | 检索结果中相关文档比例 | 相关文档数 / Top-K |
| Context Recall | 所有相关文档中被检索到的比例 | 检索到的相关 / 总相关 |
| Faithfulness | 生成内容是否忠实于上下文 | LLM-as-Judge |
| Answer Relevance | 回答与问题的相关度 | Embedding 相似度 |
| Answer Correctness | 回答的正确性 | 与标准答案比对 |
9.2 Ragas 自动评估
from ragas import evaluate
from ragas.metrics import faithfulness, answer_relevancy, context_precision
from datasets import Dataset
def evaluate_rag(questions: list[str], answers: list[str], contexts: list[list[str]]):
"""使用 Ragas 框架评估 RAG 质量。"""
data = Dataset.from_dict({
"question": questions,
"answer": answers,
"contexts": contexts,
})
result = evaluate(
dataset=data,
metrics=[faithfulness, answer_relevancy, context_precision],
)
return result.to_pandas()
9.3 RAG 优化检查清单
- Embedding 模型与数据领域匹配(通用/代码/医疗)
- 分块大小通过实验确定(512-1024 token 是常见起点)
- 混合检索(Dense + BM25)覆盖语义和关键词
- Reranker 提升 Top-K 排序质量
- Query 重写/扩展(如 HyDE:生成假答案再检索)
- 元数据过滤(如按日期、作者过滤)
- 上下文窗口不超过 LLM 限制(留出生成空间)
交叉链接:
- 向量数据库对比与集成 — PGVector/Qdrant/Chroma 深度对比
- OpenAI API 基础调用 — 流式输出与 Token 管理
- Prompt 工程与 Function Calling — RAG Prompt 模板设计
- Python 数据科学与 AI — NumPy/Embeddings 基础
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。