feat: QYLAW 法律法规知识库
- 语义检索(FAISS + embedding)+ 精确查找(法规名+条号) - RAG 问答(SSE 流式,支持 thinking 折叠显示) - 法规浏览(原文阅读) - 历史记录(检索+对话持久化到 SQLite) - 设置页(系统提示词/模板/LLM 参数可配置) - 检索质量评估脚本 Generated with [Devin](https://devin.ai) Co-Authored-By: Devin <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
@@ -0,0 +1,329 @@
|
||||
"""索引构建脚本 — 全量/增量构建 FAISS 索引 + SQLite metadata
|
||||
|
||||
用法:
|
||||
python scripts/build_index.py --mode full # 全量构建
|
||||
python scripts/build_index.py --mode incremental # 增量构建
|
||||
|
||||
流程:
|
||||
1. 扫描 law-pack 目录,按类别切片
|
||||
2. 写入 SQLite metadata(laws + clauses 表)
|
||||
3. 批量调用 embedding 服务向量化所有条文
|
||||
4. 构建 FAISS HNSW 索引并持久化
|
||||
5. 输出统计信息
|
||||
|
||||
安全:
|
||||
- 先写临时文件,成功后原子替换
|
||||
- 中断不损坏已有索引
|
||||
"""
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import sqlite3
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Any, Optional
|
||||
|
||||
# 添加项目根目录到 path
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
import numpy as np
|
||||
from app.config import (
|
||||
LAW_PACK_DIR,
|
||||
LAW_KB_DATA_DIR,
|
||||
FAISS_INDEX_PATH,
|
||||
SQLITE_PATH,
|
||||
LOGS_DIR,
|
||||
REGION_MAPPING_PATH,
|
||||
CATEGORY_DIRS,
|
||||
EMBEDDING_BATCH_SIZE,
|
||||
HNSW_M,
|
||||
HNSW_EF_CONSTRUCTION,
|
||||
)
|
||||
from scripts.parse_clause import (
|
||||
Clause,
|
||||
load_region_mapping,
|
||||
scan_category_dir,
|
||||
)
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ===== SQLite Schema =====
|
||||
SCHEMA_SQL = """
|
||||
CREATE TABLE IF NOT EXISTS laws (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
name TEXT NOT NULL,
|
||||
category TEXT NOT NULL,
|
||||
publish_date TEXT,
|
||||
province TEXT,
|
||||
city TEXT,
|
||||
region_level TEXT,
|
||||
file_path TEXT NOT NULL,
|
||||
clause_count INTEGER DEFAULT 0,
|
||||
file_mtime REAL DEFAULT 0,
|
||||
indexed_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS clauses (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
law_id INTEGER NOT NULL,
|
||||
chapter TEXT,
|
||||
clause_no TEXT,
|
||||
content TEXT NOT NULL,
|
||||
faiss_idx INTEGER,
|
||||
FOREIGN KEY (law_id) REFERENCES laws(id)
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_clauses_law_id ON clauses(law_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_clauses_faiss_idx ON clauses(faiss_idx);
|
||||
CREATE INDEX IF NOT EXISTS idx_laws_category ON laws(category);
|
||||
CREATE INDEX IF NOT EXISTS idx_laws_province ON laws(province);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS prompt_configs (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
config_key TEXT NOT NULL UNIQUE,
|
||||
config_value TEXT,
|
||||
updated_at TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_configs_key ON prompt_configs(config_key);
|
||||
"""
|
||||
|
||||
|
||||
def init_sqlite(db_path: Path):
|
||||
"""初始化 SQLite(创建表)"""
|
||||
db_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
conn = sqlite3.connect(str(db_path))
|
||||
conn.executescript(SCHEMA_SQL)
|
||||
conn.commit()
|
||||
return conn
|
||||
|
||||
|
||||
def reset_sqlite(db_path: Path):
|
||||
"""重置 SQLite(全量构建时用)"""
|
||||
if db_path.exists():
|
||||
db_path.unlink()
|
||||
return init_sqlite(db_path)
|
||||
|
||||
|
||||
async def build_index(mode: str = "full"):
|
||||
"""构建索引"""
|
||||
start_time = time.time()
|
||||
timestamp = time.strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
logger.info("=" * 60)
|
||||
logger.info(f"QYLAW 索引构建 — {mode} 模式")
|
||||
logger.info(f"时间: {timestamp}")
|
||||
logger.info(f"法规目录: {LAW_PACK_DIR}")
|
||||
logger.info(f"数据目录: {LAW_KB_DATA_DIR}")
|
||||
logger.info("=" * 60)
|
||||
|
||||
# 确保目录存在
|
||||
LAW_KB_DATA_DIR.mkdir(parents=True, exist_ok=True)
|
||||
(LAW_KB_DATA_DIR / "faiss").mkdir(parents=True, exist_ok=True)
|
||||
LOGS_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# 加载区域映射
|
||||
region_mapping = load_region_mapping(REGION_MAPPING_PATH)
|
||||
logger.info(f"加载区域映射: {len(region_mapping)} 条")
|
||||
|
||||
# 扫描所有类别
|
||||
all_clauses: List[Clause] = []
|
||||
all_skipped: List[str] = []
|
||||
category_stats: Dict[str, Dict] = {}
|
||||
|
||||
for category in CATEGORY_DIRS:
|
||||
logger.info(f"扫描 [{category}]...")
|
||||
clauses, skipped = scan_category_dir(LAW_PACK_DIR, category, region_mapping)
|
||||
all_clauses.extend(clauses)
|
||||
all_skipped.extend(skipped)
|
||||
law_count = len(set(c.law_name for c in clauses))
|
||||
category_stats[category] = {
|
||||
"laws": law_count,
|
||||
"clauses": len(clauses),
|
||||
"skipped": len(skipped),
|
||||
}
|
||||
logger.info(f" [{category}] {law_count} 法规, {len(clauses)} 条文, {len(skipped)} 跳过")
|
||||
|
||||
logger.info(f"总计: {len(all_clauses)} 条文, {len(all_skipped)} 跳过文件")
|
||||
|
||||
if not all_clauses:
|
||||
logger.error("无有效条文,终止构建")
|
||||
return
|
||||
|
||||
# 写入 SQLite
|
||||
if mode == "full":
|
||||
conn = reset_sqlite(SQLITE_PATH)
|
||||
else:
|
||||
conn = init_sqlite(SQLITE_PATH)
|
||||
|
||||
# 按法规分组
|
||||
laws_map: Dict[str, Clause] = {} # law_name -> 首个 clause(取 metadata)
|
||||
law_clauses: Dict[str, List[Clause]] = {}
|
||||
for c in all_clauses:
|
||||
if c.law_name not in laws_map:
|
||||
laws_map[c.law_name] = c
|
||||
law_clauses[c.law_name] = []
|
||||
law_clauses[c.law_name].append(c)
|
||||
|
||||
# 写 laws 表
|
||||
law_id_map: Dict[str, int] = {}
|
||||
for law_name, first_clause in laws_map.items():
|
||||
file_path = Path(first_clause.file_path)
|
||||
file_mtime = file_path.stat().st_mtime if file_path.exists() else 0
|
||||
cur = conn.execute(
|
||||
"""
|
||||
INSERT INTO laws (name, category, publish_date, province, city, region_level,
|
||||
file_path, clause_count, file_mtime, indexed_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
law_name,
|
||||
first_clause.category,
|
||||
first_clause.publish_date,
|
||||
first_clause.province,
|
||||
first_clause.city,
|
||||
first_clause.region_level,
|
||||
first_clause.file_path,
|
||||
len(law_clauses[law_name]),
|
||||
file_mtime,
|
||||
timestamp,
|
||||
),
|
||||
)
|
||||
law_id_map[law_name] = cur.lastrowid
|
||||
|
||||
# 写 clauses 表(暂不写 faiss_idx,向量化后更新)
|
||||
clause_rows: List[tuple] = []
|
||||
for law_name, clauses in law_clauses.items():
|
||||
law_id = law_id_map[law_name]
|
||||
for c in clauses:
|
||||
clause_rows.append((law_id, c.chapter, c.clause_no, c.content))
|
||||
|
||||
conn.executemany(
|
||||
"INSERT INTO clauses (law_id, chapter, clause_no, content) VALUES (?, ?, ?, ?)",
|
||||
clause_rows,
|
||||
)
|
||||
conn.commit()
|
||||
|
||||
# 获取 clause id 顺序(与写入顺序一致)
|
||||
clause_ids = [row[0] for row in conn.execute("SELECT id FROM clauses ORDER BY id").fetchall()]
|
||||
clause_texts = [row[0] for row in conn.execute("SELECT content FROM clauses ORDER BY id").fetchall()]
|
||||
logger.info(f"SQLite 写入完成: {len(clause_ids)} 条文, {len(law_id_map)} 法规")
|
||||
|
||||
# 批量向量化
|
||||
logger.info(f"开始向量化(批量大小 {EMBEDDING_BATCH_SIZE})...")
|
||||
|
||||
# 延迟导入 embedding 服务(避免循环依赖)
|
||||
from app.services.embedding import embed_batch
|
||||
|
||||
embed_start = time.time()
|
||||
try:
|
||||
all_vecs = await embed_batch(clause_texts, batch_size=EMBEDDING_BATCH_SIZE)
|
||||
except Exception as e:
|
||||
logger.error(f"向量化失败: {e}")
|
||||
conn.close()
|
||||
return
|
||||
|
||||
embed_time = time.time() - embed_start
|
||||
logger.info(f"向量化完成: {len(all_vecs)} 向量, 耗时 {embed_time:.0f}s")
|
||||
|
||||
if len(all_vecs) != len(clause_ids):
|
||||
logger.error(f"向量数 {len(all_vecs)} != 条文数 {len(clause_ids)}")
|
||||
conn.close()
|
||||
return
|
||||
|
||||
vec_dim = len(all_vecs[0]) if all_vecs else 0
|
||||
logger.info(f"向量维度: {vec_dim}")
|
||||
|
||||
# 更新 clauses 表的 faiss_idx
|
||||
for i, cid in enumerate(clause_ids):
|
||||
conn.execute("UPDATE clauses SET faiss_idx = ? WHERE id = ?", (i, cid))
|
||||
conn.commit()
|
||||
|
||||
# 构建 FAISS 索引
|
||||
# 注:HNSW 构建内存峰值高(68万向量约 30GB),114 内存紧张时 OOM
|
||||
# 改用 IndexFlatL2(暴力检索,零额外内存,68万向量检索 <100ms)
|
||||
# 若内存充裕可改回 HNSW:faiss.IndexHNSWFlat(vec_dim, HNSW_M)
|
||||
logger.info("构建 FAISS Flat 索引(暴力检索,省内存)...")
|
||||
import faiss
|
||||
|
||||
vecs_array = np.array(all_vecs, dtype=np.float32)
|
||||
index = faiss.IndexFlatL2(vec_dim)
|
||||
index.add(vecs_array)
|
||||
|
||||
logger.info(f"FAISS 索引构建完成: {index.ntotal} 向量")
|
||||
|
||||
# 原子写入(先写临时文件,成功后替换)
|
||||
tmp_faiss = FAISS_INDEX_PATH.with_suffix(".faiss.tmp")
|
||||
faiss.write_index(index, str(tmp_faiss))
|
||||
|
||||
# 替换
|
||||
if FAISS_INDEX_PATH.exists():
|
||||
FAISS_INDEX_PATH.unlink()
|
||||
tmp_faiss.rename(FAISS_INDEX_PATH)
|
||||
logger.info(f"FAISS 索引持久化: {FAISS_INDEX_PATH}")
|
||||
|
||||
conn.close()
|
||||
|
||||
# 输出统计
|
||||
total_time = time.time() - start_time
|
||||
index_size = FAISS_INDEX_PATH.stat().st_size / 1024 / 1024
|
||||
|
||||
stats = {
|
||||
"mode": mode,
|
||||
"timestamp": timestamp,
|
||||
"total_clauses": len(all_clauses),
|
||||
"total_laws": len(law_id_map),
|
||||
"category_stats": category_stats,
|
||||
"skipped_files": all_skipped,
|
||||
"vector_dim": vec_dim,
|
||||
"faiss_index_size_mb": round(index_size, 1),
|
||||
"embed_time_seconds": round(embed_time, 0),
|
||||
"total_time_seconds": round(total_time, 0),
|
||||
}
|
||||
|
||||
# 写统计到日志文件
|
||||
stats_path = LOGS_DIR / f"build_{time.strftime('%Y%m%d_%H%M%S')}.json"
|
||||
with open(stats_path, "w", encoding="utf-8") as f:
|
||||
json.dump(stats, f, ensure_ascii=False, indent=2)
|
||||
|
||||
# 控制台输出
|
||||
print("\n" + "=" * 60)
|
||||
print(f" QYLAW 索引构建完成 — {mode} 模式")
|
||||
print(f" 耗时: {total_time:.0f}s ({total_time/3600:.1f}h)")
|
||||
print("=" * 60)
|
||||
for cat, s in category_stats.items():
|
||||
print(f" [{cat}] {s['laws']} 法规, {s['clauses']} 条文, {s['skipped']} 跳过")
|
||||
print(f" 总切片: {len(all_clauses)} 条")
|
||||
print(f" 向量维度: {vec_dim}")
|
||||
print(f" FAISS 索引: {index_size:.1f} MB")
|
||||
print(f" 向量化耗时: {embed_time:.0f}s")
|
||||
print(f" 跳过文件: {len(all_skipped)} 篇")
|
||||
print(f" 索引文件: {FAISS_INDEX_PATH}")
|
||||
print(f" SQLite: {SQLITE_PATH}")
|
||||
print(f" 统计日志: {stats_path}")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="QYLAW 索引构建")
|
||||
parser.add_argument(
|
||||
"--mode",
|
||||
choices=["full", "incremental"],
|
||||
default="full",
|
||||
help="构建模式: full=全量重建, incremental=增量更新",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
asyncio.run(build_index(args.mode))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Executable
+107
@@ -0,0 +1,107 @@
|
||||
#!/bin/bash
|
||||
# QYLAW 部署脚本 — 传输数据 + 构建镜像 + 启动容器
|
||||
# 用法: bash scripts/deploy.sh
|
||||
set -e
|
||||
|
||||
# ===== 配置 =====
|
||||
REMOTE_HOST="nvidia@192.168.110.114"
|
||||
REMOTE_LAW_PACK="/data/law-pack-2026-07-01-markdown"
|
||||
REMOTE_LAW_KB_DATA="/data/law-kb-data"
|
||||
LOCAL_LAW_PACK="/Users/freedak/Documents/AIDashboard/qy123/law-pack-2026-07-01-markdown"
|
||||
PROJECT_DIR="$(cd "$(dirname "$0")/.." && pwd)"
|
||||
|
||||
echo "============================================"
|
||||
echo " QYLAW 法律法规知识库部署"
|
||||
echo " 时间: $(date '+%Y-%m-%d %H:%M:%S')"
|
||||
echo "============================================"
|
||||
|
||||
# ===== Step 1: 传输法规数据(若 114 上不存在)=====
|
||||
echo ""
|
||||
echo "[Step 1] 检查并传输法规数据..."
|
||||
ssh "$REMOTE_HOST" "test -d $REMOTE_LAW_PACK && echo EXISTS || echo MISSING"
|
||||
LAW_PACK_STATUS=$(ssh "$REMOTE_HOST" "test -d $REMOTE_LAW_PACK && echo EXISTS || echo MISSING")
|
||||
|
||||
if [ "$LAW_PACK_STATUS" = "MISSING" ]; then
|
||||
echo " 法规数据不存在,开始 rsync 传输(310MB)..."
|
||||
if [ ! -d "$LOCAL_LAW_PACK" ]; then
|
||||
echo " [错误] 本地法规数据不存在: $LOCAL_LAW_PACK"
|
||||
exit 1
|
||||
fi
|
||||
rsync -avz --progress "$LOCAL_LAW_PACK/" "$REMOTE_HOST:$REMOTE_LAW_PACK/"
|
||||
echo " [OK] 法规数据传输完成"
|
||||
else
|
||||
echo " [OK] 法规数据已存在,跳过传输"
|
||||
fi
|
||||
|
||||
# ===== Step 2: 创建索引数据目录 =====
|
||||
echo ""
|
||||
echo "[Step 2] 创建索引数据目录..."
|
||||
ssh "$REMOTE_HOST" "mkdir -p $REMOTE_LAW_KB_DATA/faiss $REMOTE_LAW_KB_DATA/logs"
|
||||
echo " [OK] 目录就绪"
|
||||
|
||||
# ===== Step 3: 传输项目代码并构建镜像 =====
|
||||
echo ""
|
||||
echo "[Step 3] 传输代码并构建 Docker 镜像..."
|
||||
ssh "$REMOTE_HOST" "mkdir -p /data/project/law-kb"
|
||||
rsync -avz --exclude '__pycache__' --exclude '.git' --exclude 'law-kb-data' \
|
||||
"$PROJECT_DIR/" "$REMOTE_HOST:/data/project/law-kb/"
|
||||
|
||||
echo " 构建 Docker 镜像..."
|
||||
ssh "$REMOTE_HOST" "cd /data/project/law-kb && docker build -t law-kb:latest ."
|
||||
echo " [OK] 镜像构建完成"
|
||||
|
||||
# ===== Step 4: 停止旧容器(若存在)=====
|
||||
echo ""
|
||||
echo "[Step 4] 停止旧容器..."
|
||||
ssh "$REMOTE_HOST" "docker rm -f law-kb 2>/dev/null || true"
|
||||
echo " [OK] 旧容器已清理"
|
||||
|
||||
# ===== Step 5: 启动新容器 =====
|
||||
echo ""
|
||||
echo "[Step 5] 启动 law-kb 容器..."
|
||||
ssh "$REMOTE_HOST" "
|
||||
docker run -d --name law-kb --network host --restart unless-stopped \
|
||||
-v $REMOTE_LAW_PACK:$REMOTE_LAW_PACK:ro \
|
||||
-v $REMOTE_LAW_KB_DATA:$REMOTE_LAW_KB_DATA \
|
||||
-e LAW_PACK_DIR=$REMOTE_LAW_PACK \
|
||||
-e LAW_KB_DATA_DIR=$REMOTE_LAW_KB_DATA \
|
||||
-e EMBEDDING_URL=http://localhost:8003/v1 \
|
||||
-e LLM_URL=http://localhost:7000/v1 \
|
||||
law-kb:latest
|
||||
"
|
||||
echo " [OK] 容器已启动"
|
||||
|
||||
# ===== Step 6: 健康检查 =====
|
||||
echo ""
|
||||
echo "[Step 6] 健康检查..."
|
||||
sleep 3
|
||||
echo " 等待服务启动..."
|
||||
for i in $(seq 1 10); do
|
||||
HEALTH=$(ssh "$REMOTE_HOST" "curl -s -o /dev/null -w '%{http_code}' http://localhost:8090/health 2>/dev/null || echo 000")
|
||||
if [ "$HEALTH" = "200" ]; then
|
||||
echo " [OK] 服务健康(尝试 $i)"
|
||||
break
|
||||
fi
|
||||
echo " 等待中...($i/10, HTTP $HEALTH)"
|
||||
sleep 2
|
||||
done
|
||||
|
||||
# 统计
|
||||
echo ""
|
||||
echo " 索引状态:"
|
||||
ssh "$REMOTE_HOST" "curl -s http://localhost:8090/api/stats 2>/dev/null | python3 -m json.tool 2>/dev/null || echo ' 统计接口未就绪(可能需要先构建索引)'"
|
||||
|
||||
echo ""
|
||||
echo "============================================"
|
||||
echo " 部署完成!"
|
||||
echo " 访问: http://192.168.110.114:8090"
|
||||
echo ""
|
||||
echo " 下一步:"
|
||||
echo " 1. 若索引未构建,执行:"
|
||||
echo " ssh $REMOTE_HOST"
|
||||
echo " docker run --rm --network host \\"
|
||||
echo " -v $REMOTE_LAW_PACK:$REMOTE_LAW_PACK:ro \\"
|
||||
echo " -v $REMOTE_LAW_KB_DATA:$REMOTE_LAW_KB_DATA \\"
|
||||
echo " law-kb:latest python scripts/build_index.py --mode full"
|
||||
echo " 2. 构建完成后重启容器: docker restart law-kb"
|
||||
echo "============================================"
|
||||
@@ -0,0 +1,348 @@
|
||||
#!/usr/bin/env python3
|
||||
"""检索质量评估脚本
|
||||
|
||||
从 metadata.db 随机抽取条款,自动生成测试 query,
|
||||
调用 /api/search 接口,计算 Recall@K、MRR、NDCG 等指标。
|
||||
|
||||
用法:
|
||||
python3 scripts/eval_retrieval.py [--sample 500] [--top-k 10] [--api http://localhost:8090]
|
||||
|
||||
两种测试模式:
|
||||
1. 结构化 query:"法规名 条号"(如"郑州市劳动用工条例 第三十二条")
|
||||
→ 测试精确查找模式(keyword)和语义检索模式(semantic)的命中差异
|
||||
2. 内容 query:条款内容前 50 字
|
||||
→ 测试语义检索是否能通过条文片段找到原文
|
||||
"""
|
||||
import argparse
|
||||
import asyncio
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import sqlite3
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Tuple
|
||||
|
||||
# 添加项目根目录到 path
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
def build_eval_dataset(db_path: str, sample_size: int = 500) -> List[Dict]:
|
||||
"""从数据库随机抽取条款,构建评估数据集
|
||||
|
||||
Returns:
|
||||
[{"clause_id", "law_name", "clause_no", "content", "law_id"}, ...]
|
||||
"""
|
||||
conn = sqlite3.connect(db_path)
|
||||
conn.row_factory = sqlite3.Row
|
||||
|
||||
# 随机抽取条款(排除内容过短的)
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT c.id, c.law_id, c.clause_no, c.content,
|
||||
l.name as law_name, l.category, l.province
|
||||
FROM clauses c JOIN laws l ON c.law_id = l.id
|
||||
WHERE length(c.content) > 20 AND c.clause_no IS NOT NULL
|
||||
ORDER BY RANDOM() LIMIT ?
|
||||
""",
|
||||
(sample_size,),
|
||||
).fetchall()
|
||||
conn.close()
|
||||
|
||||
dataset = []
|
||||
for r in rows:
|
||||
dataset.append({
|
||||
"clause_id": r["id"],
|
||||
"law_id": r["law_id"],
|
||||
"law_name": r["law_name"],
|
||||
"clause_no": r["clause_no"],
|
||||
"content": r["content"],
|
||||
"category": r["category"],
|
||||
"province": r["province"],
|
||||
})
|
||||
return dataset
|
||||
|
||||
|
||||
def generate_queries(item: Dict) -> Dict[str, str]:
|
||||
"""为一个条款生成多种 query
|
||||
|
||||
Returns:
|
||||
{"structured": "法规名 条号", "content_snippet": "条款前50字"}
|
||||
"""
|
||||
# 结构化 query:法规名 + 条号
|
||||
structured = f"{item['law_name']} {item['clause_no']}"
|
||||
|
||||
# 内容 query:条款内容前 50 字(去掉换行)
|
||||
content_clean = item["content"].replace("\n", " ").strip()
|
||||
snippet = content_clean[:50]
|
||||
|
||||
return {
|
||||
"structured": structured,
|
||||
"content_snippet": snippet,
|
||||
}
|
||||
|
||||
|
||||
async def search_api(
|
||||
client: httpx.AsyncClient,
|
||||
api_base: str,
|
||||
query: str,
|
||||
mode: str = "semantic",
|
||||
top_k: int = 10,
|
||||
) -> List[Dict]:
|
||||
"""调用检索 API
|
||||
|
||||
Returns:
|
||||
结果列表,每项含 clause_id/law_name/clause_no/score 等
|
||||
"""
|
||||
params = {
|
||||
"query": query,
|
||||
"mode": mode,
|
||||
"top_k": top_k,
|
||||
"page": 1,
|
||||
"page_size": top_k,
|
||||
}
|
||||
try:
|
||||
resp = await client.get(f"{api_base}/api/search", params=params, timeout=30)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
results = data.get("data", {}).get("results", [])
|
||||
# 补充 clause_id(检索 API 返回的字段名)
|
||||
for r in results:
|
||||
if "clause_id" not in r:
|
||||
r["clause_id"] = r.get("id")
|
||||
return results
|
||||
except Exception as e:
|
||||
print(f" [ERROR] 检索失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
def compute_recall_at_k(results: List[Dict], gt_clause_id: int, k: int) -> float:
|
||||
"""Recall@K:ground truth 是否在 Top-K 结果中"""
|
||||
top_k = results[:k]
|
||||
for r in top_k:
|
||||
if r.get("clause_id") == gt_clause_id:
|
||||
return 1.0
|
||||
return 0.0
|
||||
|
||||
|
||||
def compute_mrr(results: List[Dict], gt_clause_id: int) -> float:
|
||||
"""MRR:第一个相关文档的排名倒数"""
|
||||
for i, r in enumerate(results, 1):
|
||||
if r.get("clause_id") == gt_clause_id:
|
||||
return 1.0 / i
|
||||
return 0.0
|
||||
|
||||
|
||||
def compute_ndcg_at_k(results: List[Dict], gt_clause_id: int, k: int) -> float:
|
||||
"""NDCG@K:二值相关性(命中=1,未命中=0)"""
|
||||
dcg = 0.0
|
||||
for i, r in enumerate(results[:k], 1):
|
||||
if r.get("clause_id") == gt_clause_id:
|
||||
dcg = 1.0 / math.log2(i + 1)
|
||||
break
|
||||
# IDCG:理想情况下相关文档排第 1
|
||||
idcg = 1.0 / math.log2(2) # = 1.0
|
||||
return dcg / idcg if idcg > 0 else 0.0
|
||||
|
||||
|
||||
def compute_hit_rate(results: List[Dict], gt_clause_id: int) -> float:
|
||||
"""Hit Rate:是否至少命中一个相关文档"""
|
||||
for r in results:
|
||||
if r.get("clause_id") == gt_clause_id:
|
||||
return 1.0
|
||||
return 0.0
|
||||
|
||||
|
||||
async def run_eval(
|
||||
dataset: List[Dict],
|
||||
api_base: str,
|
||||
top_k: int = 10,
|
||||
modes: List[str] = None,
|
||||
query_types: List[str] = None,
|
||||
) -> Dict:
|
||||
"""运行评估
|
||||
|
||||
Args:
|
||||
dataset: 评估数据集
|
||||
api_base: API 地址
|
||||
top_k: Top-K
|
||||
modes: 检索模式列表 ["semantic", "keyword"]
|
||||
query_types: query 类型列表 ["structured", "content_snippet"]
|
||||
|
||||
Returns:
|
||||
评估结果报告
|
||||
"""
|
||||
if modes is None:
|
||||
modes = ["semantic", "keyword"]
|
||||
if query_types is None:
|
||||
query_types = ["structured", "content_snippet"]
|
||||
|
||||
report = {
|
||||
"sample_size": len(dataset),
|
||||
"top_k": top_k,
|
||||
"modes": modes,
|
||||
"query_types": query_types,
|
||||
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"results": {},
|
||||
}
|
||||
|
||||
async with httpx.AsyncClient() as client:
|
||||
for mode in modes:
|
||||
for qtype in query_types:
|
||||
key = f"{mode}_{qtype}"
|
||||
print(f"\n=== 评估: mode={mode}, query_type={qtype} ===")
|
||||
|
||||
recalls = []
|
||||
mrrs = []
|
||||
ndcgs = []
|
||||
hit_rates = []
|
||||
latencies = []
|
||||
errors = 0
|
||||
|
||||
for i, item in enumerate(dataset):
|
||||
queries = generate_queries(item)
|
||||
query = queries.get(qtype, "")
|
||||
if not query:
|
||||
continue
|
||||
|
||||
gt_id = item["clause_id"]
|
||||
|
||||
t0 = time.time()
|
||||
results = await search_api(client, api_base, query, mode, top_k)
|
||||
latency = (time.time() - t0) * 1000
|
||||
latencies.append(latency)
|
||||
|
||||
if not results:
|
||||
errors += 1
|
||||
recalls.append(0.0)
|
||||
mrrs.append(0.0)
|
||||
ndcgs.append(0.0)
|
||||
hit_rates.append(0.0)
|
||||
else:
|
||||
recalls.append(compute_recall_at_k(results, gt_id, top_k))
|
||||
mrrs.append(compute_mrr(results, gt_id))
|
||||
ndcgs.append(compute_ndcg_at_k(results, gt_id, top_k))
|
||||
hit_rates.append(compute_hit_rate(results, gt_id))
|
||||
|
||||
# 进度
|
||||
if (i + 1) % 50 == 0:
|
||||
avg_recall = sum(recalls) / len(recalls)
|
||||
print(f" 进度: {i+1}/{len(dataset)} | Recall@{top_k}={avg_recall:.3f}")
|
||||
|
||||
n = len(recalls)
|
||||
report["results"][key] = {
|
||||
"mode": mode,
|
||||
"query_type": qtype,
|
||||
"count": n,
|
||||
"errors": errors,
|
||||
"recall_at_k": sum(recalls) / n if n else 0,
|
||||
"mrr": sum(mrrs) / n if n else 0,
|
||||
"ndcg_at_k": sum(ndcgs) / n if n else 0,
|
||||
"hit_rate": sum(hit_rates) / n if n else 0,
|
||||
"avg_latency_ms": sum(latencies) / len(latencies) if latencies else 0,
|
||||
"p95_latency_ms": sorted(latencies)[int(len(latencies) * 0.95)] if latencies else 0,
|
||||
}
|
||||
|
||||
r = report["results"][key]
|
||||
print(f" 结果: Recall@{top_k}={r['recall_at_k']:.3f} | MRR={r['mrr']:.3f} | "
|
||||
f"NDCG@{top_k}={r['ndcg_at_k']:.3f} | HitRate={r['hit_rate']:.3f} | "
|
||||
f"avg={r['avg_latency_ms']:.0f}ms | errors={errors}")
|
||||
|
||||
return report
|
||||
|
||||
|
||||
def print_report(report: Dict):
|
||||
"""打印评估报告"""
|
||||
print("\n" + "=" * 80)
|
||||
print("检索质量评估报告")
|
||||
print("=" * 80)
|
||||
print(f"时间: {report['timestamp']}")
|
||||
print(f"样本数: {report['sample_size']}")
|
||||
print(f"Top-K: {report['top_k']}")
|
||||
print()
|
||||
|
||||
# 表格输出
|
||||
print(f"{'模式':<12} {'Query类型':<18} {'Recall@K':>10} {'MRR':>10} {'NDCG@K':>10} {'HitRate':>10} {'avg(ms)':>10} {'错误':>6}")
|
||||
print("-" * 96)
|
||||
for key, r in report["results"].items():
|
||||
print(f"{r['mode']:<12} {r['query_type']:<18} "
|
||||
f"{r['recall_at_k']:>10.3f} {r['mrr']:>10.3f} {r['ndcg_at_k']:>10.3f} "
|
||||
f"{r['hit_rate']:>10.3f} {r['avg_latency_ms']:>10.0f} {r['errors']:>6}")
|
||||
|
||||
print()
|
||||
# 分析建议
|
||||
sem_struct = report["results"].get("semantic_structured", {})
|
||||
kw_struct = report["results"].get("keyword_structured", {})
|
||||
sem_content = report["results"].get("semantic_content_snippet", {})
|
||||
|
||||
print("分析:")
|
||||
if sem_struct and kw_struct:
|
||||
if kw_struct["recall_at_k"] > sem_struct["recall_at_k"] + 0.1:
|
||||
print(f" - 结构化查询(法规名+条号):精确查找(Recall={kw_struct['recall_at_k']:.3f})"
|
||||
f" 显著优于语义检索(Recall={sem_struct['recall_at_k']:.3f})")
|
||||
print(f" → 建议:用户输入法规名+条号时,自动切换到精确查找模式")
|
||||
else:
|
||||
print(f" - 结构化查询:语义检索(Recall={sem_struct['recall_at_k']:.3f})"
|
||||
f" 与精确查找(Recall={kw_struct['recall_at_k']:.3f})相当")
|
||||
|
||||
if sem_content:
|
||||
r = sem_content
|
||||
if r["recall_at_k"] > 0.8:
|
||||
print(f" - 内容片段查询:语义检索表现优秀(Recall={r['recall_at_k']:.3f})")
|
||||
elif r["recall_at_k"] > 0.5:
|
||||
print(f" - 内容片段查询:语义检索表现一般(Recall={r['recall_at_k']:.3f}),有提升空间")
|
||||
else:
|
||||
print(f" - 内容片段查询:语义检索表现较差(Recall={r['recall_at_k']:.3f}),需检查 embedding 模型")
|
||||
|
||||
print()
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="检索质量评估")
|
||||
parser.add_argument("--sample", type=int, default=500, help="抽样数量(默认 500)")
|
||||
parser.add_argument("--top-k", type=int, default=10, help="Top-K(默认 10)")
|
||||
parser.add_argument("--api", type=str, default="http://localhost:8090", help="API 地址")
|
||||
parser.add_argument("--db", type=str, default="/data/law-kb-data/metadata.db", help="SQLite 路径")
|
||||
parser.add_argument("--seed", type=int, default=42, help="随机种子(默认 42)")
|
||||
parser.add_argument("--output", type=str, default=None, help="结果输出 JSON 文件路径")
|
||||
parser.add_argument("--modes", type=str, nargs="+", default=["semantic", "keyword"],
|
||||
help="检索模式(默认 semantic keyword)")
|
||||
parser.add_argument("--query-types", type=str, nargs="+",
|
||||
default=["structured", "content_snippet"],
|
||||
help="query 类型(默认 structured content_snippet)")
|
||||
args = parser.parse_args()
|
||||
|
||||
random.seed(args.seed)
|
||||
|
||||
print(f"构建评估数据集(抽样 {args.sample} 条)...")
|
||||
dataset = build_eval_dataset(args.db, args.sample)
|
||||
print(f"数据集大小: {len(dataset)}")
|
||||
|
||||
if not dataset:
|
||||
print("错误:数据集为空,检查数据库路径")
|
||||
sys.exit(1)
|
||||
|
||||
# 运行评估
|
||||
report = asyncio.run(run_eval(
|
||||
dataset=dataset,
|
||||
api_base=args.api,
|
||||
top_k=args.top_k,
|
||||
modes=args.modes,
|
||||
query_types=args.query_types,
|
||||
))
|
||||
|
||||
# 打印报告
|
||||
print_report(report)
|
||||
|
||||
# 保存结果
|
||||
if args.output:
|
||||
with open(args.output, "w", encoding="utf-8") as f:
|
||||
json.dump(report, f, ensure_ascii=False, indent=2)
|
||||
print(f"结果已保存到: {args.output}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,266 @@
|
||||
"""法规切片解析器 — 按"第X条"切片,提取 metadata
|
||||
|
||||
解析 markdown 法规文件:
|
||||
- 法规名:文件名去 .md 后缀,去末尾 _YYYYMMDD
|
||||
- 发布日期:文件名末尾 _YYYYMMDD
|
||||
- 章节:跟踪"第X章"上下文
|
||||
- 条文:按"第X条"切片,内容聚合到下一个"第X条"/"第X章"/"第X节"
|
||||
- 地方性法规:从 地方性法规区域映射.json 附加省份/市级
|
||||
|
||||
用法:
|
||||
from scripts.parse_clause import parse_law_file
|
||||
clauses = parse_law_file("/data/law-pack/法律/中华人民共和国民法典_20200528.md", "法律")
|
||||
"""
|
||||
import json
|
||||
import re
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import List, Optional, Dict, Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 中文数字正则(支持"第一""第二十二""第一百二十六"等)
|
||||
CN_NUM = r"[一二三四五六七八九十百千零〇两]+"
|
||||
# 条文开头:第X条 + 全角空格/普通空格
|
||||
CLAUSE_RE = re.compile(rf"^第({CN_NUM})条[\s\u3000]+(.*)")
|
||||
# 章节开头:第X章 + 标题
|
||||
CHAPTER_RE = re.compile(rf"^第{CN_NUM}章[\s\u3000]+(.+)")
|
||||
# 节开头:第X节
|
||||
SECTION_RE = re.compile(rf"^第{CN_NUM}节[\s\u3000]+(.+)")
|
||||
# 编开头:第X编
|
||||
PART_RE = re.compile(rf"^第{CN_NUM}编[\s\u3000]+(.+)")
|
||||
# 文件名日期后缀
|
||||
DATE_SUFFIX_RE = re.compile(r"_(\d{8})$")
|
||||
|
||||
|
||||
@dataclass
|
||||
class Clause:
|
||||
"""法规条文"""
|
||||
law_name: str
|
||||
category: str
|
||||
chapter: Optional[str] = None # 当前章节,如"第一章 基本规定"
|
||||
clause_no: str = "" # 条号,如"第一条" / "第143条"
|
||||
content: str = "" # 条文内容
|
||||
publish_date: Optional[str] = None # 发布日期 YYYY-MM-DD
|
||||
file_path: str = "" # 原文路径
|
||||
province: Optional[str] = None # 省份(地方性法规)
|
||||
city: Optional[str] = None # 市级(地方性法规)
|
||||
region_level: Optional[str] = None # 区域级别(省级/市级)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"law_name": self.law_name,
|
||||
"category": self.category,
|
||||
"chapter": self.chapter,
|
||||
"clause_no": self.clause_no,
|
||||
"content": self.content,
|
||||
"publish_date": self.publish_date,
|
||||
"file_path": self.file_path,
|
||||
"province": self.province,
|
||||
"city": self.city,
|
||||
"region_level": self.region_level,
|
||||
}
|
||||
|
||||
|
||||
def extract_name_and_date(filename: str) -> tuple[str, Optional[str]]:
|
||||
"""从文件名提取法规名和发布日期
|
||||
|
||||
Args:
|
||||
filename: 文件名(含或不含 .md 后缀)
|
||||
|
||||
Returns:
|
||||
(法规名, 发布日期 YYYY-MM-DD 或 None)
|
||||
"""
|
||||
stem = Path(filename).stem
|
||||
match = DATE_SUFFIX_RE.search(stem)
|
||||
if match:
|
||||
date_str = match.group(1)
|
||||
name = stem[: match.start()]
|
||||
# 格式化日期 YYYY-MM-DD
|
||||
date_formatted = f"{date_str[:4]}-{date_str[4:6]}-{date_str[6:8]}"
|
||||
return name, date_formatted
|
||||
return stem, None
|
||||
|
||||
|
||||
def load_region_mapping(mapping_path: Path) -> Dict[str, Optional[Dict]]:
|
||||
"""加载地方性法规区域映射
|
||||
|
||||
Returns:
|
||||
{文件名: {region_name, province, level, region_id, province_id} | None}
|
||||
"""
|
||||
if not mapping_path.exists():
|
||||
logger.warning(f"区域映射文件不存在: {mapping_path}")
|
||||
return {}
|
||||
with open(mapping_path, "r", encoding="utf-8") as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def parse_law_file(
|
||||
filepath: str | Path,
|
||||
category: str,
|
||||
region_mapping: Optional[Dict] = None,
|
||||
) -> List[Clause]:
|
||||
"""解析单个法规文件,按"第X条"切片
|
||||
|
||||
Args:
|
||||
filepath: 法规 markdown 文件路径
|
||||
category: 法规类别(法律/行政法规/监察法规/司法解释/地方性法规)
|
||||
region_mapping: 地方性法规区域映射(文件名 -> {province, city, ...})
|
||||
|
||||
Returns:
|
||||
条文列表;若文件无法识别"第X条"则返回空列表
|
||||
"""
|
||||
filepath = Path(filepath)
|
||||
filename = filepath.name
|
||||
law_name, publish_date = extract_name_and_date(filename)
|
||||
|
||||
# 地方性法规附加区域信息
|
||||
province = None
|
||||
city = None
|
||||
region_level = None
|
||||
if category == "地方性法规" and region_mapping is not None:
|
||||
info = region_mapping.get(filename)
|
||||
if info is None:
|
||||
# 映射为 null,跳过(调用方负责记录)
|
||||
return []
|
||||
province = info.get("province")
|
||||
region_level = info.get("level")
|
||||
# 市级:region_name 如果不是省份名,则视为市级
|
||||
region_name = info.get("region_name")
|
||||
if region_name and province and region_name != province:
|
||||
city = region_name
|
||||
|
||||
# 读取文件
|
||||
try:
|
||||
text = filepath.read_text(encoding="utf-8")
|
||||
except Exception as e:
|
||||
logger.error(f"读取失败 {filepath}: {e}")
|
||||
return []
|
||||
|
||||
lines = text.split("\n")
|
||||
clauses: List[Clause] = []
|
||||
current_chapter: Optional[str] = None
|
||||
current_clause: Optional[Clause] = None
|
||||
current_content_lines: List[str] = []
|
||||
|
||||
def _flush_current():
|
||||
"""将当前条文写入列表"""
|
||||
nonlocal current_clause, current_content_lines
|
||||
if current_clause is not None:
|
||||
content = "\n".join(current_content_lines).strip()
|
||||
if content:
|
||||
current_clause.content = content
|
||||
current_clause.chapter = current_chapter
|
||||
clauses.append(current_clause)
|
||||
current_clause = None
|
||||
current_content_lines = []
|
||||
|
||||
for line in lines:
|
||||
stripped = line.strip()
|
||||
|
||||
# 跳过空行(但保留在内容中,后续 strip 处理)
|
||||
if not stripped:
|
||||
if current_clause is not None:
|
||||
current_content_lines.append("")
|
||||
continue
|
||||
|
||||
# 检测章节
|
||||
chap_match = CHAPTER_RE.match(stripped)
|
||||
if chap_match:
|
||||
_flush_current()
|
||||
current_chapter = stripped
|
||||
continue
|
||||
|
||||
# 检测节(不切片,只更新上下文)
|
||||
if SECTION_RE.match(stripped):
|
||||
_flush_current()
|
||||
continue
|
||||
|
||||
# 检测编(不切片)
|
||||
if PART_RE.match(stripped):
|
||||
_flush_current()
|
||||
continue
|
||||
|
||||
# 检测条文
|
||||
clause_match = CLAUSE_RE.match(stripped)
|
||||
if clause_match:
|
||||
_flush_current()
|
||||
clause_no = f"第{clause_match.group(1)}条"
|
||||
first_line = clause_match.group(2).strip()
|
||||
current_clause = Clause(
|
||||
law_name=law_name,
|
||||
category=category,
|
||||
clause_no=clause_no,
|
||||
publish_date=publish_date,
|
||||
file_path=str(filepath),
|
||||
province=province,
|
||||
city=city,
|
||||
region_level=region_level,
|
||||
)
|
||||
current_content_lines = [first_line] if first_line else []
|
||||
continue
|
||||
|
||||
# 普通行:若在条文中,追加到内容
|
||||
if current_clause is not None:
|
||||
current_content_lines.append(stripped)
|
||||
|
||||
_flush_current()
|
||||
return clauses
|
||||
|
||||
|
||||
def scan_category_dir(
|
||||
law_pack_dir: Path,
|
||||
category: str,
|
||||
region_mapping: Optional[Dict] = None,
|
||||
) -> tuple[List[Clause], List[str]]:
|
||||
"""扫描某类别目录下所有法规文件
|
||||
|
||||
Args:
|
||||
law_pack_dir: 法规根目录
|
||||
category: 类别(法律/行政法规/...)
|
||||
region_mapping: 区域映射(仅地方性法规需要)
|
||||
|
||||
Returns:
|
||||
(条文列表, 跳过文件列表[null 映射或解析失败])
|
||||
"""
|
||||
cat_dir = law_pack_dir / category
|
||||
if not cat_dir.exists():
|
||||
logger.warning(f"类别目录不存在: {cat_dir}")
|
||||
return [], []
|
||||
|
||||
all_clauses: List[Clause] = []
|
||||
skipped: List[str] = []
|
||||
|
||||
md_files = sorted(cat_dir.glob("*.md"))
|
||||
for md_file in md_files:
|
||||
clauses = parse_law_file(md_file, category, region_mapping)
|
||||
if not clauses:
|
||||
# 区分 null 映射跳过 vs 解析失败
|
||||
if category == "地方性法规" and region_mapping is not None:
|
||||
info = region_mapping.get(md_file.name)
|
||||
if info is None:
|
||||
skipped.append(f"[null映射] {md_file.name}")
|
||||
else:
|
||||
skipped.append(f"[无条文] {md_file.name}")
|
||||
else:
|
||||
skipped.append(f"[无条文] {md_file.name}")
|
||||
else:
|
||||
all_clauses.extend(clauses)
|
||||
|
||||
return all_clauses, skipped
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# 自测:解析民法典
|
||||
import sys
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
test_file = sys.argv[1] if len(sys.argv) > 1 else None
|
||||
if test_file:
|
||||
clauses = parse_law_file(test_file, "法律")
|
||||
print(f"解析 {test_file}: {len(clauses)} 条")
|
||||
for c in clauses[:3]:
|
||||
print(f" {c.clause_no} [{c.chapter}] {c.content[:50]}...")
|
||||
else:
|
||||
print("用法: python parse_clause.py <法规文件路径>")
|
||||
@@ -0,0 +1,147 @@
|
||||
"""FAISS 索引重建脚本 — 从 SQLite 读取已有条文,重新向量化 + 构建 FAISS
|
||||
|
||||
用途:索引构建因 OOM/中断失败,但 SQLite metadata 已完好时,
|
||||
跳过切片解析,仅重新向量化 + 构建 FAISS 索引。
|
||||
|
||||
用法:
|
||||
python scripts/rebuild_faiss.py
|
||||
"""
|
||||
import asyncio
|
||||
import logging
|
||||
import sqlite3
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
import numpy as np
|
||||
from app.config import (
|
||||
SQLITE_PATH,
|
||||
FAISS_INDEX_PATH,
|
||||
LAW_KB_DATA_DIR,
|
||||
EMBEDDING_BATCH_SIZE,
|
||||
)
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(message)s",
|
||||
)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def rebuild():
|
||||
"""从 SQLite 读取条文,重新向量化 + 构建 FAISS Flat 索引"""
|
||||
start_time = time.time()
|
||||
timestamp = time.strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
logger.info("=" * 60)
|
||||
logger.info(f"QYLAW FAISS 索引重建(跳过切片,从 SQLite 读取)")
|
||||
logger.info(f"时间: {timestamp}")
|
||||
logger.info("=" * 60)
|
||||
|
||||
if not SQLITE_PATH.exists():
|
||||
logger.error(f"SQLite 不存在: {SQLITE_PATH}")
|
||||
return
|
||||
|
||||
conn = sqlite3.connect(str(SQLITE_PATH))
|
||||
conn.row_factory = sqlite3.Row
|
||||
|
||||
# 读取所有条文(按 faiss_idx 排序,保证顺序一致)
|
||||
rows = conn.execute(
|
||||
"SELECT id, faiss_idx, content FROM clauses WHERE faiss_idx IS NOT NULL ORDER BY faiss_idx"
|
||||
).fetchall()
|
||||
total = len(rows)
|
||||
logger.info(f"从 SQLite 读取: {total} 条文")
|
||||
|
||||
if total == 0:
|
||||
logger.error("无条文,终止")
|
||||
conn.close()
|
||||
return
|
||||
|
||||
clause_texts = [row["content"] for row in rows]
|
||||
clause_ids = [row["id"] for row in rows]
|
||||
|
||||
# 向量化(分批写入预分配的 numpy 数组,避免 Python list 内存爆炸)
|
||||
logger.info(f"开始向量化(批量大小 {EMBEDDING_BATCH_SIZE})...")
|
||||
from app.services.embedding import embed_batch
|
||||
|
||||
embed_start = time.time()
|
||||
|
||||
# 先用第一批获取向量维度,然后预分配 numpy 数组
|
||||
first_batch = clause_texts[:EMBEDDING_BATCH_SIZE]
|
||||
try:
|
||||
first_vecs = await embed_batch(first_batch, batch_size=EMBEDDING_BATCH_SIZE)
|
||||
except Exception as e:
|
||||
logger.error(f"向量化失败: {e}")
|
||||
conn.close()
|
||||
return
|
||||
|
||||
vec_dim = len(first_vecs[0])
|
||||
logger.info(f"向量维度: {vec_dim}")
|
||||
|
||||
# 预分配 numpy 数组(68万 × 1024 × 4字节 ≈ 2.6GB,一次性分配)
|
||||
vecs_array = np.zeros((total, vec_dim), dtype=np.float32)
|
||||
vecs_array[:len(first_vecs)] = np.array(first_vecs, dtype=np.float32)
|
||||
logger.info(f"预分配 numpy 数组: {total} × {vec_dim} ({vecs_array.nbytes / 1024**3:.1f} GB)")
|
||||
|
||||
# 分批向量化剩余条文,直接写入 numpy 数组
|
||||
for i in range(EMBEDDING_BATCH_SIZE, total, EMBEDDING_BATCH_SIZE):
|
||||
batch = clause_texts[i : i + EMBEDDING_BATCH_SIZE]
|
||||
try:
|
||||
batch_vecs = await embed_batch(batch, batch_size=EMBEDDING_BATCH_SIZE)
|
||||
vecs_array[i : i + len(batch_vecs)] = np.array(batch_vecs, dtype=np.float32)
|
||||
except Exception as e:
|
||||
logger.error(f"向量化批次 {i} 失败: {e}")
|
||||
conn.close()
|
||||
return
|
||||
if (i // EMBEDDING_BATCH_SIZE) % 1000 == 0:
|
||||
logger.info(f"进度: {i}/{total} ({i*100//total}%)")
|
||||
|
||||
embed_time = time.time() - embed_start
|
||||
logger.info(f"向量化完成: {total} 向量, 耗时 {embed_time:.0f}s")
|
||||
|
||||
# 释放 clause_texts 内存(不再需要)
|
||||
del clause_texts
|
||||
|
||||
# 构建 FAISS Flat 索引(省内存,暴力检索)
|
||||
logger.info("构建 FAISS Flat 索引...")
|
||||
import faiss
|
||||
|
||||
index = faiss.IndexFlatL2(vec_dim)
|
||||
index.add(vecs_array)
|
||||
logger.info(f"FAISS 索引构建完成: {index.ntotal} 向量")
|
||||
|
||||
# 释放 numpy 数组(已写入 FAISS 索引)
|
||||
del vecs_array
|
||||
|
||||
# 原子写入
|
||||
FAISS_INDEX_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
tmp_faiss = FAISS_INDEX_PATH.with_suffix(".faiss.tmp")
|
||||
faiss.write_index(index, str(tmp_faiss))
|
||||
if FAISS_INDEX_PATH.exists():
|
||||
FAISS_INDEX_PATH.unlink()
|
||||
tmp_faiss.rename(FAISS_INDEX_PATH)
|
||||
|
||||
index_size = FAISS_INDEX_PATH.stat().st_size / 1024 / 1024
|
||||
logger.info(f"FAISS 索引持久化: {FAISS_INDEX_PATH} ({index_size:.1f} MB)")
|
||||
|
||||
conn.close()
|
||||
|
||||
total_time = time.time() - start_time
|
||||
print("\n" + "=" * 60)
|
||||
print(f" FAISS 索引重建完成")
|
||||
print(f" 总耗时: {total_time:.0f}s ({total_time/60:.1f}min)")
|
||||
print(f" 向量数: {total}")
|
||||
print(f" 向量维度: {vec_dim}")
|
||||
print(f" 索引大小: {index_size:.1f} MB")
|
||||
print(f" 向量化耗时: {embed_time:.0f}s")
|
||||
print("=" * 60)
|
||||
|
||||
|
||||
def main():
|
||||
asyncio.run(rebuild())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,101 @@
|
||||
"""索引验证脚本 — 预置典型查询,抽样验证检索质量
|
||||
|
||||
用法:
|
||||
python scripts/verify_index.py
|
||||
|
||||
输出每个查询的 Top-5 结果,供人工评估命中率。
|
||||
"""
|
||||
import asyncio
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent))
|
||||
|
||||
from app.services.retriever import Retriever
|
||||
from app.services.embedding import embed_text
|
||||
|
||||
|
||||
# 预置典型查询(20 个,覆盖各类别)
|
||||
TEST_QUERIES = [
|
||||
# 法律
|
||||
"个人信息保护",
|
||||
"个人信息跨境传输",
|
||||
"竞业协议补偿金",
|
||||
"民法典婚姻家庭",
|
||||
"合同违约责任",
|
||||
"知识产权侵权赔偿",
|
||||
"公司股东权利",
|
||||
# 行政法规
|
||||
"不动产登记流程",
|
||||
"专利申请条件",
|
||||
# 司法解释
|
||||
"公益诉讼办案规则",
|
||||
"刑事诉讼证据规则",
|
||||
# 监察法规
|
||||
"监察工作信息公开",
|
||||
# 地方性法规
|
||||
"垃圾分类管理",
|
||||
"烟花爆竹禁放",
|
||||
"物业管理规定",
|
||||
"生态环境保护",
|
||||
"城市市容管理",
|
||||
"食品安全监管",
|
||||
"道路交通管理",
|
||||
"未成年人保护",
|
||||
]
|
||||
|
||||
|
||||
async def verify():
|
||||
"""执行验证"""
|
||||
# 加载索引
|
||||
retriever = Retriever.get_instance()
|
||||
retriever.load()
|
||||
|
||||
if not retriever.is_ready():
|
||||
print("索引未加载,请先执行构建")
|
||||
return
|
||||
|
||||
print("=" * 70)
|
||||
print(" QYLAW 索引验证")
|
||||
print("=" * 70)
|
||||
|
||||
hit_count = 0
|
||||
total = len(TEST_QUERIES)
|
||||
|
||||
for i, query in enumerate(TEST_QUERIES, 1):
|
||||
print(f"\n[{i}/{total}] 查询: {query}")
|
||||
try:
|
||||
query_vec = await embed_text(query)
|
||||
results = retriever.search(query_vec, top_k=5)
|
||||
if results:
|
||||
print(f" Top-5 结果:")
|
||||
for j, r in enumerate(results, 1):
|
||||
score = r["score"]
|
||||
law = r["law_name"]
|
||||
clause = r.get("clause_no", "")
|
||||
chapter = r.get("chapter", "")
|
||||
content_preview = r["content"][:80].replace("\n", " ")
|
||||
print(f" {j}. [{score:.4f}] {law} {clause} ({chapter})")
|
||||
print(f" {content_preview}...")
|
||||
# 简单命中率判断:Top-1 分数 > 0.5 视为命中
|
||||
if results[0]["score"] > 0.5:
|
||||
hit_count += 1
|
||||
print(f" ✅ 命中")
|
||||
else:
|
||||
print(f" ⚠️ 分数偏低")
|
||||
else:
|
||||
print(f" ❌ 无结果")
|
||||
except Exception as e:
|
||||
print(f" ❌ 错误: {e}")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print(f" 验证完成: {hit_count}/{total} 命中 (Top-1 score > 0.5)")
|
||||
print("=" * 70)
|
||||
|
||||
|
||||
def main():
|
||||
asyncio.run(verify())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user