rag_kb · RAG + ChromaDB 检索管线¶
本页收录可运行示例,配套本文档「RAG 检索管线」(RAG 检索管线)。 代码为随框架交付的示例,已一并收录到本站供直接对照使用。
说明与运行¶
示例: 知识库(RAG)检索 —— CSV 文档 → 切分/向量化 → ChromaDB 存储 → 检索. 演示 agentframework.retrieval 能力包对接真实第三方向量库 ChromaDB 的全流程: 1. 切分: 读取《知识库文档.csv》(公积金办事指南问答, 每行一问答, 每行视为一个文档), 长文本经 RecursiveCharacterSplitter 递归切分; 2. 向量化: embedder 将每个块编码为向量(默认 HashEmbedder 离线可跑, 语义检索请替换 _build_embedder 为真实中文 embedding 实现); 3. 向量存储: ChromaVectorStore 实现框架 VectorStore 协议, 持久化到 chromadb.PersistentClient(默认 examples/rag/.chroma_db); 4. 检索: create_retriever + search_with_citations 输出带引用 (source/chunk_index/score) 的命中, 支持交互提问. 运行: cd agentframework
首次: 建 venv(需 python>=3.11)并安装框架依赖 + 示例依赖 chromadb¶
python3 -m venv .venv && source .venv/bin/activate pip install -e ".[dev]" chromadb
之后(在已安装依赖的 venv 内执行; 缺 pydantic/chromadb 说明当前解释器未装依赖)¶
PYTHONPATH=src python examples/rag/rag_kb.py # 幂等重建 + 演示查询(交互输入 Ctrl-C 退出) PYTHONPATH=src python examples/rag/rag_kb.py --query "租房提取住房公积金办理材料" PYTHONPATH=src python examples/rag/rag_kb.py --no-ingest --query "..." # 只查不重建 PYTHONPATH=src python examples/rag/rag_kb.py --fresh # 更换 embedder / 维度变更时清空重建
运行前需在已安装
agentframework的 Python ≥ 3.11 环境;示例多为自带本地模型,无需真实 API Key;需要外部依赖(如chromadb/ 数据库驱动)会在说明中注明。
完整源码¶
"""示例: 知识库(RAG)检索 —— CSV 文档 → 切分/向量化 → ChromaDB 存储 → 检索.
演示 agentframework.retrieval 能力包对接真实第三方向量库 ChromaDB 的全流程:
1. 切分: 读取《知识库文档.csv》(公积金办事指南问答, 每行一问答, 每行视为一个文档),
长文本经 RecursiveCharacterSplitter 递归切分;
2. 向量化: embedder 将每个块编码为向量(默认 HashEmbedder 离线可跑,
语义检索请替换 _build_embedder 为真实中文 embedding 实现);
3. 向量存储: ChromaVectorStore 实现框架 VectorStore 协议, 持久化到
chromadb.PersistentClient(默认 examples/rag/.chroma_db);
4. 检索: create_retriever + search_with_citations 输出带引用
(source/chunk_index/score) 的命中, 支持交互提问.
运行:
cd agentframework
# 首次: 建 venv(需 python>=3.11)并安装框架依赖 + 示例依赖 chromadb
python3 -m venv .venv && source .venv/bin/activate
pip install -e ".[dev]" chromadb
# 之后(在已安装依赖的 venv 内执行; 缺 pydantic/chromadb 说明当前解释器未装依赖)
PYTHONPATH=src python examples/rag/rag_kb.py # 幂等重建 + 演示查询(交互输入 Ctrl-C 退出)
PYTHONPATH=src python examples/rag/rag_kb.py --query "租房提取住房公积金办理材料"
PYTHONPATH=src python examples/rag/rag_kb.py --no-ingest --query "..." # 只查不重建
PYTHONPATH=src python examples/rag/rag_kb.py --fresh # 更换 embedder / 维度变更时清空重建
"""
from __future__ import annotations
import argparse
import csv
import hashlib
import math
import os
import sys
from pathlib import Path
from agentframework.retrieval import (
BaseEmbedder,
DocumentChunk,
HashEmbedder,
RecursiveCharacterSplitter,
SearchResult,
VectorStore,
create_retriever,
ingest_document,
)
# 演示查询: 前三条取自知识库原文标题(HashEmbedder 精确匹配命中),
# 最后一条为自然口语改写(近似匹配分数较低, 语义 embedding 下效果更佳).
DEMO_QUERIES = [
"租房提取住房公积金办理材料",
"提前偿还住房公积金贷款办理材料",
"商业性个人住房贷款转住房公积金贷款的申请条件",
"离职很久了想一次性取公积金需要什么?",
]
# 默认 chroma 持久化目录(与 csv 同目录, 已 gitignore), collection 名。
DEFAULT_CHROMA_DIR = Path(__file__).resolve().parent / ".chroma_db"
DEFAULT_COLLECTION = "zhijin_kb"
class ChromaVectorStore:
"""ChromaDB 向量存储(实现框架 retrieval.VectorStore 协议).
框架 VectorStore 只约定 add/search/delete_by_source/count 四个方法
(见 src/agentframework/retrieval/store.py), 存哪个向量库由使用方决定;
本例对接 chromadb 并落盘持久化。要点:
- add/search 前对向量做 L2 归一化, 使 chroma 的 L2 距离与余弦相似度同序,
返回分数换算: 余弦相似度 = 1 - d^2 / 2 (0~1, 越大越相关, 与 MemoryVectorStore 对齐);
- 块元数据只写 str/int/bool(chroma 约束), id 用 source_index_hash 保证稳定唯一。
"""
def __init__(self, path: str, collection_name: str = DEFAULT_COLLECTION):
"""创建持久化客户端并获取/创建 collection.
Args:
path: chroma 落盘目录(不存在会自动创建).
collection_name: collection 名(更换 embedder/维度时用 drop_collection 重建).
Raises:
RuntimeError: chromadb 未安装时给出安装指引.
"""
try:
import chromadb
from chromadb.config import Settings
except ImportError as exc: # pragma: no cover - 依赖缺失提示
raise RuntimeError("示例依赖 chromadb 未安装, 请先执行: pip install chromadb") from exc
os.makedirs(path, exist_ok=True)
self._name = collection_name
self._client = chromadb.PersistentClient(
path=path, settings=Settings(anonymized_telemetry=False)
)
self._coll = self._get_or_create_collection()
def _get_or_create_collection(self):
"""取已有 collection; 不存在则创建(显式 l2 空间, 配合向量归一化做余弦)."""
try:
return self._client.get_collection(name=self._name)
except Exception: # chroma 缺 collection 时抛异常, 捕获后创建
return self._client.create_collection(name=self._name, metadata={"hnsw:space": "l2"})
# ---------- VectorStore 协议实现 ----------
def add(self, chunks: list[DocumentChunk], vectors: list[list[float]]) -> None:
"""批量写入块 + 归一化向量到 chroma."""
if len(chunks) != len(vectors):
raise ValueError("chunks 与 vectors 数量不一致")
if not chunks:
return
ids: list[str] = []
documents: list[str] = []
metadatas: list[dict] = []
embeddings: list[list[float]] = []
for index, (chunk, vector) in enumerate(zip(chunks, vectors, strict=True)):
source = chunk.metadata.get("source", f"chunk-{index}")
chunk_index = chunk.metadata.get("chunk_index", index)
ids.append(_stable_id(source, chunk_index, chunk.text))
documents.append(chunk.text)
metadatas.append(dict(chunk.metadata))
embeddings.append(_normalize_vector(vector))
self._coll.add(ids=ids, embeddings=embeddings, documents=documents, metadatas=metadatas)
def search(self, vector: list[float], *, top_k: int = 5) -> list[SearchResult]:
"""按归一化向量在 chroma 检索最相似的块(余弦相似度排序)."""
if top_k <= 0:
return []
total = self._coll.count()
if total == 0:
return []
query_vector = _normalize_vector(vector)
n_results = min(top_k, total)
result = self._coll.query(query_embeddings=[query_vector], n_results=n_results)
results: list[SearchResult] = []
documents = (result.get("documents") or [[]])[0] or []
metadatas = (result.get("metadatas") or [[]])[0] or []
distances = (result.get("distances") or [[]])[0] or []
for text, meta, distance in zip(documents, metadatas, distances, strict=True):
if text is None or meta is None:
continue
# 单位向量下: ||u-v||^2 = 2 - 2cos => cos = 1 - d^2/2
score = 1.0 - (float(distance) ** 2) / 2.0
score = max(0.0, min(1.0, score))
results.append(
SearchResult(chunk=DocumentChunk(text=text, metadata=dict(meta)), score=score)
)
return results
def delete_by_source(self, source: str) -> int:
"""按来源删除该文档全部块(重导入时先删旧), 返回删除条数."""
found = self._coll.get(where={"source": source}, include=[])
ids = found.get("ids") or []
if ids:
self._coll.delete(ids=ids)
return len(ids)
def count(self) -> int:
"""当前块总数."""
return self._coll.count()
def drop_collection(self) -> None:
"""删除并重建 collection(--fresh / embedder 维度变更用)."""
try:
self._client.delete_collection(name=self._name)
except Exception: # 不存在时忽略
pass
self._coll = self._get_or_create_collection()
def _normalize_vector(vector: list[float]) -> list[float]:
"""L2 归一化(使 L2 距离序等价于余弦相似度序)."""
norm = math.sqrt(sum(v * v for v in vector))
if norm == 0:
return list(vector)
return [v / norm for v in vector]
def _stable_id(source: str, chunk_index: int, text: str) -> str:
"""生成稳定唯一 id: {source}_{chunk_index}_{text sha256 前 12 位}."""
digest = hashlib.sha256(text.encode("utf-8")).hexdigest()[:12]
return f"{source}_{chunk_index}_{digest}"
def _build_embedder() -> BaseEmbedder:
"""构建 embedding 实现(默认离线确定性哈希, 维度 64).
HashEmbedder 仅适合链路演示(相同文本→相同向量, 精确/词重叠匹配);
生产/语义检索请替换为真实中文 embedding, 例如 OpenAI 兼容 API:
class OpenAIEmbedder(BaseEmbedder):
def __init__(self, api_key, base_url, model, dimension):
self._client = httpx.Client(base_url=base_url,
headers={"Authorization": f"Bearer {api_key}"})
self._model, self._dimension = model, dimension
@property
def dimension(self) -> int:
return self._dimension
def embed(self, texts):
resp = self._client.post("/embeddings",
json={"model": self._model, "input": list(texts)})
return [d["embedding"] for d in resp.json()["data"]]
注意: 更换 embedder 后向量维度可能变化, 需用 --fresh 清空重建 collection.
"""
return HashEmbedder(dimension=64)
def load_kb_csv(path: Path) -> list[dict]:
"""解析《知识库文档.csv》为问答文档列表(每行一个文档).
Args:
path: csv 文件路径(utf-8-sig, 每行至多 2 列: 标题/问题, 答复/材料).
Returns:
[{index, title, source, text}] 列表, 已跳过空行; text = "标题\\n答复".
Raises:
FileNotFoundError: 文件不存在.
RuntimeError: 无有效问答行.
"""
if not path.exists():
raise FileNotFoundError(f"知识库文件不存在: {path}")
docs: list[dict] = []
with open(path, encoding="utf-8-sig", newline="") as fh:
for row_number, row in enumerate(csv.reader(fh), start=1):
title = (row[0] if row else "").strip()
body = (row[1] if len(row) > 1 else "").strip()
if not title and not body:
continue
text = f"{title}\n{body}" if body else title
docs.append(
{
"index": row_number,
"title": title,
"source": f"row-{row_number}",
"text": text,
}
)
if not docs:
raise RuntimeError(f"知识库文件无有效内容: {path}")
return docs
def build_kb(
store: VectorStore,
embedder: BaseEmbedder,
docs: list[dict],
*,
chunk_size: int,
chunk_overlap: int,
) -> tuple[int, int]:
"""逐条把问答文档导入 chroma(幂等: 同 source 重导入先删旧块).
Args:
store: 向量存储(ChromaVectorStore 实例).
embedder: embedding 实现.
docs: load_kb_csv 返回的问答文档列表.
chunk_size / chunk_overlap: RecursiveCharacterSplitter 切分参数.
Returns:
(入库总块数, 被切成多块的长文档数).
"""
splitter = RecursiveCharacterSplitter(chunk_size=chunk_size, chunk_overlap=chunk_overlap)
total_chunks = 0
multi_chunk_docs = 0
for doc in docs:
chunks = ingest_document(
doc["text"],
store=store,
embedder=embedder,
source=doc["source"],
splitter=splitter,
metadata={"title": doc["title"]},
)
total_chunks += len(chunks)
if len(chunks) > 1:
multi_chunk_docs += 1
return total_chunks, multi_chunk_docs
def print_citations(citations: list[dict]) -> None:
"""打印检索命中的引用结构 [{source, chunk_index, text, score}]."""
if not citations:
print(" 未检索到相关内容(若为空库请先不带 --no-ingest 运行入库).")
return
print(f" 命中 {len(citations)} 条:")
for i, citation in enumerate(citations, start=1):
source = citation["source"]
chunk_index = citation["chunk_index"]
score = citation["score"]
text = citation["text"].strip().replace("\n", " / ")
if len(text) > 120:
text = text[:120] + "…"
print(f" [{i}] {source}#{chunk_index} score={score:.4f}")
print(f" {text}")
def run_queries(retriever, queries: list[str], *, top_k: int) -> None:
"""对一组查询执行检索并打印引用结果."""
for query in queries:
print(f"\n查询: {query}")
citations = retriever.search_with_citations(query, top_k=top_k)
print_citations(citations)
def _read_line() -> str | None:
"""从标准输入读一行, 兼容非 UTF-8 终端粘贴(如 GBK), 避免 input() 抛 UnicodeDecodeError.
优先按 UTF-8 解码; 失败则退回 GB18030(覆盖 GBK/GB2312); 仍失败则以替换符收尾
(不中断交互)。注意: 建议终端保持 UTF-8 编码以获得最佳检索效果。
Returns:
去掉换行的行文本; 读到 EOF(Ctrl-D)返回 None.
"""
buffer = getattr(sys.stdin, "buffer", None)
if buffer is None: # 非缓冲 stdin 时退回 input()
try:
line = input("")
except EOFError:
return None
return line.rstrip("\r\n")
raw = buffer.readline()
if raw == b"": # EOF
return None
for encoding in ("utf-8", "gb18030"):
try:
return raw.decode(encoding).rstrip("\r\n")
except UnicodeDecodeError:
continue
return raw.decode("utf-8", errors="replace").rstrip("\r\n")
def run_interactive(retriever, *, top_k: int) -> None:
"""交互式提问(仅 tty; Ctrl-C / EOF / exit 退出)."""
print("\n交互问答(输入 exit/quit 或 Ctrl-C 退出):")
try:
while True:
line = _read_line()
if line is None: # EOF(Ctrl-D)
print()
break
query = line.strip()
if not query:
continue
if query.lower() in {"exit", "quit"}:
break
citations = retriever.search_with_citations(query, top_k=top_k)
print_citations(citations)
except KeyboardInterrupt:
print("\n再见。")
def _parse_args(argv: list[str] | None) -> argparse.Namespace:
"""解析命令行参数."""
default_csv = Path(__file__).resolve().parent / "知识库文档.csv"
parser = argparse.ArgumentParser(
description="知识库(RAG)检索示例: CSV → 切分/向量化 → ChromaDB → 检索"
)
parser.add_argument("--csv", type=Path, default=default_csv, help="知识库问答 CSV")
parser.add_argument(
"--chroma-dir", type=str, default=str(DEFAULT_CHROMA_DIR), help="chroma 落盘目录"
)
parser.add_argument("--collection", default=DEFAULT_COLLECTION, help="chroma collection 名")
parser.add_argument("--no-ingest", action="store_true", help="跳过重建, 只查现有库")
parser.add_argument(
"--fresh", action="store_true", help="清空重建 collection(更换 embedder/维度变更用)"
)
parser.add_argument("--chunk-size", type=int, default=200, help="切分块大小(默认 200)")
parser.add_argument("--chunk-overlap", type=int, default=50, help="切分块重叠(默认 50)")
parser.add_argument("--top-k", type=int, default=3, help="检索返回条数(默认 3)")
parser.add_argument("--query", help="单次检索问题(缺省跑演示查询, tty 下进入交互)")
return parser.parse_args(argv)
def main(argv: list[str] | None = None) -> int:
"""示例入口: 建库(可选) → 检索 → 交互."""
args = _parse_args(argv)
store = ChromaVectorStore(args.chroma_dir, collection_name=args.collection)
if args.fresh:
store.drop_collection()
embedder = _build_embedder()
if not args.no_ingest:
docs = load_kb_csv(args.csv)
print(f"解析知识库: {len(docs)} 行问答 <- {args.csv}")
try:
total_chunks, multi_chunk_docs = build_kb(
store,
embedder,
docs,
chunk_size=args.chunk_size,
chunk_overlap=args.chunk_overlap,
)
except Exception as exc: # 向 CLI 用户转友好提示
message = str(exc)
if "dimension" in message.lower():
print("入库失败: 向量维度与现有 collection 不一致(可能更换了 embedder)。")
print("提示: 请加 --fresh 清空重建后重试。")
else:
print(f"入库失败: {message}")
return 1
print(
f"入库完成: {total_chunks} 个块"
f"({multi_chunk_docs} 行长文本被切为多块), 向量维度={embedder.dimension}"
)
else:
total_chunks = store.count()
print(f"跳过重建, 读取现有库: {total_chunks} 个块, collection={args.collection!r}")
if total_chunks == 0:
print("提示: 知识库为空, 请先不带 --no-ingest 运行一次完成入库。")
return 0
retriever = create_retriever(store, embedder, top_k=args.top_k)
if args.query:
print(f"\n===== 检索: {args.query} =====")
run_queries(retriever, [args.query], top_k=args.top_k)
return 0
print("\n===== 演示查询 =====")
run_queries(retriever, DEMO_QUERIES, top_k=args.top_k)
if sys.stdin.isatty():
run_interactive(retriever, top_k=args.top_k)
return 0
if __name__ == "__main__":
raise SystemExit(main())
回到示例索引。