-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathvector_store.py
More file actions
233 lines (186 loc) · 7.45 KB
/
Copy pathvector_store.py
File metadata and controls
233 lines (186 loc) · 7.45 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
"""
向量存储模块: 把 Embedding 后的向量 + 原文 + 元数据存入 ChromaDB
ChromaDB 是嵌入式向量数据库(类似 SQLite, 不需要单独启动服务器)
- Collection: 相当于数据库里的一张表
- add(): 往表里插数据(向量 + 原文 + 元数据)
- query(): 用向量搜最相似的结果
【面试重点】向量数据库的作用:
- 持久化存储向量, 程序重启不丢失
- 高效的近似最近邻搜索(ANN), 不需要跟每个向量逐一比较
"""
import os
import chromadb
import numpy as np
from document_loader import Document
from embedder import embed_documents
# ============ 配置 ============
# ChromaDB 数据存储目录, 所有向量和元数据都持久化在这里
# 相当于 SQLite 的 .db 文件存放位置
DEFAULT_DB_DIR = "./chroma_data"
# Collection 名称, 相当于数据库里的表名
DEFAULT_COLLECTION_NAME = "documents"
# ============ 初始化 ChromaDB ============
def get_collection(
db_dir: str = DEFAULT_DB_DIR,
collection_name: str = DEFAULT_COLLECTION_NAME,
) -> chromadb.Collection:
"""
获取(或创建) ChromaDB 的 Collection.
参数:
db_dir: 数据库文件存储目录
collection_name: Collection 名称(表名)
返回:
chromadb.Collection 对象, 后续用它来 add/query
ChromaDB 两种模式:
- chromadb.Client(): 纯内存, 程序关了就没了(适合测试)
- chromadb.PersistentClient(path=...): 持久化到磁盘(我们用这个)
"""
# PersistentClient: 数据存到磁盘, 重启后还在
# 相当于: 打开(或创建)一个数据库文件
client = chromadb.PersistentClient(path=db_dir)
# get_or_create_collection: 有就打开, 没有就新建
# 相当于 SQL: CREATE TABLE IF NOT EXISTS documents
#
# metadata={"hnsw:space": "cosine"}: 告诉 ChromaDB 用余弦距离衡量相似度
#
# 【面试重点】HNSW = Hierarchical Navigable Small World
# 一种高效的近似最近邻(ANN)算法, 不用跟每个向量逐一比较
# 时间复杂度: 暴力搜索 O(n) → HNSW O(log n)
collection = client.get_or_create_collection(
name=collection_name,
metadata={"hnsw:space": "cosine"},
)
return collection
# ============ 生成唯一 ID ============
def make_chunk_id(filepath: str, chunk_index: int) -> str:
"""
为每个 chunk 生成全局唯一的 ID.
用文件名 + chunk 序号组合, 保证不同文件的 chunk 不会冲突.
例: "mcp.md_chunk_0", "mcp.md_chunk_1", "rag.md_chunk_0"
参数:
filepath: 原文件路径(从 metadata 里取)
chunk_index: 第几个分块(从 0 开始)
返回:
唯一 ID 字符串
"""
# os.path.basename(): 从完整路径里取文件名
# "F:/docs/mcp.md" → "mcp.md"
filename = os.path.basename(filepath)
return f"{filename}_chunk_{chunk_index}"
# ============ 核心函数: 存入文档 ============
def add_documents(
chunks: list[Document],
collection: chromadb.Collection = None,
) -> int:
"""
把分块后的文档列表存入 ChromaDB.
完整流程:
1. 给每个 chunk 生成唯一 ID
2. 调用 embed_documents() 算向量
3. 调用 collection.upsert() 存入 ChromaDB
参数:
chunks: split_documents() 输出的 Document 列表
collection: ChromaDB Collection, 不传则用默认的
返回:
成功存入的文档数量
【面试重点】upsert vs add:
- add(): 如果 ID 已存在会报错
- upsert(): ID 存在就更新, 不存在就插入(update + insert)
- 我们用 upsert, 这样重复索引同一个文件不会出错
"""
if not chunks:
print("没有文档需要存储")
return 0
if collection is None:
collection = get_collection()
# --- 第一步: 生成 ID ---
ids = []
for chunk in chunks:
chunk_id = make_chunk_id(
filepath=chunk.metadata.get("filepath", "unknown"),
chunk_index=chunk.metadata.get("chunk_index", 0),
)
ids.append(chunk_id)
# --- 第二步: 算向量 ---
print(f"正在计算 {len(chunks)} 个分块的 Embedding...")
vectors = embed_documents(chunks)
# --- 第三步: 存入 ChromaDB ---
# ChromaDB 要求的数据格式:
# ids: list[str] → 每条数据的唯一标识
# embeddings: list[list] → 向量(ChromaDB 要 Python list, 不能直接传 numpy)
# documents: list[str] → 原文(搜到后返回给用户看)
# metadatas: list[dict] → 元数据(溯源信息)
#
# numpy 的 .tolist() 方法: 把 numpy 数组转成 Python 嵌套列表
# 因为 ChromaDB 内部用 JSON 序列化, numpy 数组不能直接序列化
# 准备元数据: ChromaDB 的 metadata value 只支持 str/int/float/bool
# 不能存 list 或嵌套 dict, 所以我们只取需要的字段
metadatas = []
for chunk in chunks:
metadatas.append({
"filename": str(chunk.metadata.get("filename", "")),
"filepath": str(chunk.metadata.get("filepath", "")),
"chunk_index": int(chunk.metadata.get("chunk_index", 0)),
"chunk_total": int(chunk.metadata.get("chunk_total", 0)),
})
collection.upsert(
ids=ids,
embeddings=vectors.tolist(),
documents=[chunk.content for chunk in chunks],
metadatas=metadatas,
)
print(f"存储完成: {len(chunks)} 个分块已存入 ChromaDB")
return len(chunks)
# ============ 辅助函数: 查看存储状态 ============
def get_store_stats(collection: chromadb.Collection = None) -> dict:
"""
查看当前 Collection 的状态信息.
返回:
包含统计信息的字典
"""
if collection is None:
collection = get_collection()
count = collection.count() # 当前存了多少条数据
return {
"total_chunks": count,
"collection_name": collection.name,
}
# ============ 辅助函数: 清空存储 ============
def clear_store(db_dir: str = DEFAULT_DB_DIR) -> None:
"""
删除整个 ChromaDB 数据目录, 从头开始.
谨慎使用, 数据删了就没了.
"""
import shutil
if os.path.exists(db_dir):
shutil.rmtree(db_dir)
print(f"已清空存储: {db_dir}")
else:
print(f"存储目录不存在: {db_dir}")
# ============ 测试代码 ============
if __name__ == "__main__":
from document_loader import load_documents
from text_splitter import split_documents
# 先清空旧数据, 确保测试干净
clear_store()
# 完整管道: 加载 → 分块 → 存储
print("=== 第一步: 加载文档 ===")
docs = load_documents("F:/AI_Program/enterprise-kb")
print("\n=== 第二步: 分块 ===")
chunks = split_documents(docs, chunk_size=500, chunk_overlap=100)
print("\n=== 第三步: 存入 ChromaDB ===")
count = add_documents(chunks)
print("\n=== 存储状态 ===")
stats = get_store_stats()
print(f"总分块数: {stats['total_chunks']}")
print(f"Collection: {stats['collection_name']}")
# 验证: 从 ChromaDB 取出来看看
print("\n=== 验证: 查看前 3 条数据 ===")
collection = get_collection()
# peek() 随机取几条看看, 确认数据确实存进去了
sample = collection.peek(limit=3)
for i in range(min(3, len(sample["ids"]))):
print(f" ID: {sample['ids'][i]}")
print(f" 元数据: {sample['metadatas'][i]}")
print(f" 原文预览: {sample['documents'][i][:80]}...")
print()