Files
group_xinghuo_jinrong/scripts/kb/build_collections.py
T

114 lines
3.7 KiB
Python
Raw Normal View History

"""知识库入库脚本:解析桌面知识库文件 → 向量化 → 存入 Milvus。
用法:
python scripts/kb/build_collections.py # 仅创建不存在的 collection
python scripts/kb/build_collections.py --rebuild # 先删除再重建
"""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
# 把项目根目录加入 sys.path
_ROOT = Path(__file__).resolve().parent.parent.parent
sys.path.insert(0, str(_ROOT))
from app.config.settings import settings
from app.tool.document_parser import parse_collection_dir
from app.tool.embedding_tool import get_embedder
from app.tool.milvus_tool import get_milvus_client
# Collection → 子目录映射
KB_LAYOUT: dict[str, str] = {
"fin_faq": "fin_faq_collection",
"fin_product": "fin_product_collection",
"fin_policy": "fin_policy_collection",
}
# 每个 collection 的标量字段(从 chunk.metadata 中取值)
COLLECTION_FIELDS: dict[str, list[str]] = {
"fin_faq": ["question", "answer", "category", "source_doc", "chunk_no", "chunk_text"],
"fin_product": ["product_name", "risk_level", "doc_type", "source_doc", "chunk_no", "chunk_text"],
"fin_policy": ["policy_name", "chapter", "source_doc", "chunk_no", "chunk_text"],
}
def chunk_to_record(chunk, collection: str, vector: list[float]) -> dict:
"""把 Chunk + 向量转为 Milvus 插入记录。"""
record = {
"id": chunk.chunk_id,
"embedding": vector,
}
for field in COLLECTION_FIELDS[collection]:
record[field] = chunk.metadata.get(field, "")
# 类型修正:chunk_no 必须是 int
record["chunk_no"] = int(chunk.metadata.get("chunk_no", 0))
return record
def main(rebuild: bool = False) -> None:
embedder = get_embedder()
milvus = get_milvus_client()
kb_root = Path(settings.kb_root_dir)
if not kb_root.exists():
print(f"[ERROR] 知识库目录不存在: {kb_root}")
sys.exit(1)
total_inserted = 0
for coll_name, sub_dir in KB_LAYOUT.items():
dir_path = kb_root / sub_dir
if not dir_path.exists():
print(f"[WARN] 跳过 {coll_name}: 目录不存在 {dir_path}")
continue
# 1. 解析文档
print(f"\n{'=' * 60}")
print(f"[{coll_name}] 开始解析: {dir_path}")
chunks, _ = parse_collection_dir(dir_path)
print(f"[{coll_name}] 解析完成: {len(chunks)} 块")
if not chunks:
print(f"[WARN] {coll_name} 无有效内容,跳过")
continue
# 2. 重建或确保 collection
if rebuild:
milvus.drop_collection(coll_name)
milvus.ensure_collection(coll_name)
# 3. 向量化
print(f"[{coll_name}] 开始向量化 ({len(chunks)} 条)...")
texts = [c.text for c in chunks]
vectors = embedder.embed_batch(texts)
# 4. 组装记录并插入
records = []
skipped = 0
for chunk, vec in zip(chunks, vectors):
if not vec:
skipped += 1
continue
records.append(chunk_to_record(chunk, coll_name, vec))
if skipped:
print(f"[{coll_name}] 跳过 {skipped} 条(向量化失败)")
insert_count = milvus.insert(coll_name, records)
print(f"[{coll_name}] 插入 {insert_count} 条记录")
total_inserted += insert_count
print(f"\n{'=' * 60}")
print(f"入库完成:共插入 {total_inserted} 条记录到 {len(KB_LAYOUT)} 个 collection")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="知识库入库")
parser.add_argument("--rebuild", action="store_true", help="先删除再重建 collection")
args = parser.parse_args()
main(rebuild=args.rebuild)