"""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()