"""把 knowledge/_chunks.jsonl 灌入 Milvus 三个知识集合(临时脚本,跑完即删)。 设计要点: 1. **集合 schema** 按方案 §4.2 的统一字段(10 项),另加 5 个检索与合规必需的字段: `chapter`/`section`(定位)、`source_file`(溯源)、`doc_no`(内部文件编号)、 `visibility`(反洗钱手册标 internal,客服侧按此过滤)。维度 1024 与 `qwen3.7-text-embedding-flash` 实测一致,索引按方案 §2.4.4:IVF_FLAT + COSINE + nlist=128。 2. **向量输入用「标题 + 正文」**而不是只喂正文:标题里带条款号与章节名(如 「第五章 … 第十三条 … 第一类:资金流转异常」),是比正文更干净的检索信号。 3. **幂等**:用 upsert,同一 doc_id 重复灌库不会产生重复行,可以反复重跑。 4. 灌完立刻做检索自检(拿几个真实客户问题去查),不看自检结果不算灌成功。 """ import asyncio import json import os from collections import defaultdict from pathlib import Path import httpx from dotenv import load_dotenv from pymilvus import DataType, MilvusClient load_dotenv(override=False) MILVUS_URI = os.environ.get("MILVUS_URI", "http://127.0.0.1:19530") MILVUS_TOKEN = os.environ.get("MILVUS_TOKEN") or None EMBED_BASE = "https://dashscope.aliyuncs.com/compatible-mode/v1" EMBED_MODEL = "qwen3.7-text-embedding-flash" EMBED_KEY = os.environ.get("QWEN_EMBEDDING_API_KEY", "") DIM = 1024 BATCH = 10 COLLECTIONS = ["fin_faq_collection", "fin_product_collection", "fin_policy_collection"] # 检索自检用例:(自然语言问题, 期望命中的 doc_id 前缀) CHECKS = [ ("基金赎回到账需要多长时间", "FAQ-0016"), ("C1 保守型客户可以买哪些风险等级的产品", "POL-AST"), ("南方季季盈90天的起投金额是多少", "PROD"), ("开户需要准备哪些材料", "FAQ-0021"), ("业绩比较基准是什么意思", "FAQ-0019"), ("高净值客户能享受什么费率优惠", "HNW"), ] def build_schema(client: MilvusClient) -> object: schema = MilvusClient.create_schema(auto_id=False, enable_dynamic_field=False) schema.add_field("doc_id", DataType.VARCHAR, max_length=64, is_primary=True) schema.add_field("title", DataType.VARCHAR, max_length=1024) schema.add_field("content", DataType.VARCHAR, max_length=16384) schema.add_field("chapter", DataType.VARCHAR, max_length=512) schema.add_field("section", DataType.VARCHAR, max_length=512) schema.add_field("tags", DataType.VARCHAR, max_length=512) schema.add_field("doc_no", DataType.VARCHAR, max_length=64) schema.add_field("version", DataType.VARCHAR, max_length=32) schema.add_field("effective_date", DataType.VARCHAR, max_length=32) schema.add_field("expire_date", DataType.VARCHAR, max_length=32) schema.add_field("source_url", DataType.VARCHAR, max_length=512) schema.add_field("reviewer", DataType.VARCHAR, max_length=64) schema.add_field("source_file", DataType.VARCHAR, max_length=128) schema.add_field("visibility", DataType.VARCHAR, max_length=16) schema.add_field("embedding", DataType.FLOAT_VECTOR, dim=DIM) return schema def ensure_collections(client: MilvusClient) -> None: existing = set(client.list_collections()) for name in COLLECTIONS: if name in existing: print(f" 集合已存在,跳过创建:{name}") continue index_params = client.prepare_index_params() index_params.add_index( field_name="embedding", index_type="IVF_FLAT", metric_type="COSINE", params={"nlist": 128}, ) client.create_collection( collection_name=name, schema=build_schema(client), index_params=index_params ) print(f" 已创建集合:{name}") async def embed(texts: list[str]) -> list[list[float]]: async with httpx.AsyncClient(timeout=90) as client: response = await client.post( f"{EMBED_BASE}/embeddings", headers={"Authorization": f"Bearer {EMBED_KEY}"}, json={"model": EMBED_MODEL, "input": texts}, ) response.raise_for_status() payload = response.json() items = sorted(payload["data"], key=lambda item: item["index"]) return [item["embedding"] for item in items] async def main() -> None: if not EMBED_KEY: print("缺少 QWEN_EMBEDDING_API_KEY") return records = [ json.loads(line) for line in (Path("knowledge") / "_chunks.jsonl").read_text(encoding="utf-8").splitlines() if line.strip() ] print(f"待入库块数:{len(records)}") client = MilvusClient(uri=MILVUS_URI, token=MILVUS_TOKEN) print("\n== 建集合 ==") ensure_collections(client) grouped: dict[str, list[dict[str, object]]] = defaultdict(list) for record in records: grouped[str(record["collection"])].append(record) print("\n== 生成向量并写入 ==") for name, group in grouped.items(): rows: list[dict[str, object]] = [] for start in range(0, len(group), BATCH): batch = group[start:start + BATCH] vectors = await embed( [f"{record['title']}\n{record['content']}" for record in batch] ) for record, vector in zip(batch, vectors, strict=True): rows.append({ "doc_id": record["doc_id"], "title": str(record["title"])[:500], "content": str(record["content"])[:8000], "chapter": str(record["chapter"])[:250], "section": str(record["section"])[:250], "tags": str(record["tags"])[:250], "doc_no": str(record["doc_no"])[:60], "version": str(record["version"])[:30], "effective_date": str(record["effective_date"])[:30], "expire_date": str(record["expire_date"])[:30], "source_url": str(record["source_url"])[:500], "reviewer": str(record["reviewer"])[:60], "source_file": str(record["source_file"])[:120], "visibility": str(record["visibility"])[:16], "embedding": vector, }) print(f" {name}: 已向量化 {min(start + BATCH, len(group))}/{len(group)}") client.upsert(collection_name=name, data=rows) client.flush(collection_name=name) print("\n== 各集合条目数 ==") for name in COLLECTIONS: stats = client.get_collection_stats(collection_name=name) print(f" {name}: {stats.get('row_count')} 条") print("\n== 检索自检 ==") queries = [question for question, _ in CHECKS] query_vectors = await embed(queries) hits = 0 for (question, expected), vector in zip(CHECKS, query_vectors, strict=True): results = client.search( collection_name="fin_policy_collection", data=[vector], limit=3, output_fields=["doc_id", "title", "visibility"], ) top = client.search( collection_name="fin_faq_collection", data=[vector], limit=3, output_fields=["doc_id", "title"], ) product = client.search( collection_name="fin_product_collection", data=[vector], limit=3, output_fields=["doc_id", "title"], ) merged = [item for group in (top, product, results) for item in group[0]] merged.sort(key=lambda item: item["distance"], reverse=True) best = merged[0] if merged else None # 注意:COSINE 下 pymilvus 返回的 distance 越大越相似 found = "" if best is not None: found = f"{best['entity']['doc_id']} (score={best['distance']:.3f}) {best['entity']['title'][:44]}" mark = "OK " if best is not None and str(best["entity"]["doc_id"]).startswith(expected) else "检查" hits += 1 if mark == "OK " else 0 print(f" [{mark}] {question}\n → {found}") print(f"\n自检命中 {hits}/{len(CHECKS)}") asyncio.run(main())