memoryweave/scripts/reindex_bge_m3.py

148 lines
5.1 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python3
"""存量记忆向量重编码:将 LanceDB 中所有 memory vector 替换为本地 bge-m3 编码
策略:全量读 → 分批编码 → 删表重建一次写入1130 条 batch 插入)
"""
import os, sys, json, time, argparse
from typing import List
import lancedb
import pyarrow as pa
import requests
import numpy as np
EMBED_URL = os.environ.get("EMBED_URL", "http://localhost:8000/v1/embeddings")
LANCE_DIR = os.environ.get("LANCE_DIR", "/var/lib/memoryweave/lancedb")
ENCODE_BATCH = 3 # 每批编码 3 条,避开 sentence-transformers batch bug
def encode_batch(texts: List[str]) -> List[List[float]]:
"""调用本地 bge-m3 server 批量编码server 自带 prefix"""
resp = requests.post(EMBED_URL, json={"input": texts, "model": "BAAI/bge-m3"},
timeout=60)
resp.raise_for_status()
data = resp.json()["data"]
vecs = [None] * len(texts)
for item in data:
vecs[item["index"]] = item["embedding"]
return vecs
def main():
parser = argparse.ArgumentParser(description="存量记忆向量重编码")
parser.add_argument("--backup", action="store_true", help="备份当前 LanceDB")
parser.add_argument("--dry-run", action="store_true", help="只读+编码,不写入")
parser.add_argument("--batch", type=int, default=ENCODE_BATCH, help="编码批次大小")
parser.add_argument("--embed-url", default=EMBED_URL)
args = parser.parse_args()
# 1. 验证嵌入服务
try:
r = requests.get(args.embed_url.replace("/v1/embeddings", "/health"), timeout=5)
r.raise_for_status()
status = r.json()
print(f"[health] model={status.get('model','?')} device={status.get('device','?')}")
except Exception as e:
print(f"[FATAL] 嵌入服务不可用: {e}")
sys.exit(1)
# 2. 打开 LanceDB读取全量数据
db = lancedb.connect(LANCE_DIR)
tbl = db.open_table("memories")
arrow = tbl.to_arrow()
total = arrow.num_rows
print(f"[lance] 读取 {total} 条记录")
col_names = arrow.schema.names
print(f"[schema] 列: {col_names}")
# 3. 备份
if args.backup and not args.dry_run:
import shutil
bak_dir = f"{LANCE_DIR}.bak.{int(time.time())}"
print(f"[backup] 备份到 {bak_dir} ...")
shutil.copytree(LANCE_DIR, bak_dir)
print(f"[backup] 完成")
# 4. 读取 content
ids = arrow.column("id").to_pylist()
contents = arrow.column("content").to_pylist()
print(f"[data] 首条: {repr(contents[0][:60]) if contents else 'EMPTY'}")
# 5. 分批编码
print(f"[encode] 分批编码 {total} 条 (batch_size={args.batch}) ...")
all_vecs = [None] * total
errors = 0
t0 = time.time()
for offset in range(0, total, args.batch):
batch_texts = contents[offset:offset + args.batch]
try:
vecs = encode_batch(batch_texts)
for i, v in enumerate(vecs):
all_vecs[offset + i] = v
except Exception as e:
# 单条回退
for i in range(len(batch_texts)):
try:
v = encode_batch([batch_texts[i]])
all_vecs[offset + i] = v[0]
except Exception as e2:
print(f" [ERROR] id={ids[offset+i]}: {e2}")
errors += 1
elapsed = time.time() - t0
done = min(offset + args.batch, total)
rate = done / elapsed if elapsed > 0 else 0
print(f" {done}/{total} ({rate:.1f}/s) errors={errors}", end="\r")
print()
elapsed = time.time() - t0
success = total - errors
print(f"[encode] {elapsed:.1f}s | {success} 成功, {errors} 失败")
if errors > 0:
# 用零向量填充失败项
print("[encode] 用零向量填充失败项")
zero = [0.0] * 1024
for i in range(total):
if all_vecs[i] is None:
all_vecs[i] = zero
if args.dry_run:
print("[dry-run] 完成,未写入 LanceDB")
return
# 6. 建新 vector 列PyArrow FixedShapeTensorArray from numpy
vec_array = np.array(all_vecs, dtype=np.float32) # (total, 1024)
vec_col = pa.FixedShapeTensorArray.from_numpy_ndarray(vec_array)
# 替换原有 vector 列
col_index = col_names.index("vector")
new_columns = []
for i, name in enumerate(col_names):
if name == "vector":
new_columns.append(vec_col)
else:
new_columns.append(arrow.column(name))
new_arrow = pa.table(dict(zip(col_names, new_columns)))
# 7. 删表重建
print("[write] 删除旧表 ...")
db.drop_table("memories")
print("[write] 重建 memories 表 ...")
db.create_table("memories", new_arrow)
print("[write] 完成 ✓")
# 8. 验证
verify_tbl = db.open_table("memories")
verify_count = verify_tbl.count_rows()
verify_vec = verify_tbl.to_arrow().column("vector")[0]
verify_dim = len(verify_vec) if verify_vec else 0
print(f"[verify] 新表: {verify_count} 条, 向量维度={verify_dim}")
print(f"[done] 总耗时 {time.time()-t0:.1f}s")
if __name__ == "__main__":
main()