Files
ai-g3sb-backman2.0/scripts/build_fewshot_chroma_index.py
T

101 lines
3.1 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""
将 Few-shot 经验数据集(JSONL)写入 Chroma 持久化向量库。
用法(仓库根目录):
python scripts/build_fewshot_chroma_index.py
python scripts/build_fewshot_chroma_index.py --samples data/experiences/all_samples.jsonl --force
构建完成后在 .env 中设置:
FEWSHOT_USE_CHROMA=true
FEWSHOT_CHROMA_PATH=./data/embeddings/chroma_fewshot
Embedding 与 Schema 向量一致,由 USE_LOCAL_EMBEDDING / OPENAI_* / MODELSCOPE_* 等决定。
"""
from __future__ import annotations
import argparse
import json
import logging
import os
import sys
from pathlib import Path
_REPO_ROOT = Path(__file__).resolve().parents[1]
_BACKEND = _REPO_ROOT / "backend"
if str(_BACKEND) not in sys.path:
sys.path.insert(0, str(_BACKEND))
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
logger = logging.getLogger("build_fewshot_chroma")
def main() -> int:
parser = argparse.ArgumentParser(description="Few-shot JSONL → Chroma 向量索引")
parser.add_argument(
"--samples",
default=os.getenv("FEWSHOT_DATA_PATH", "data/experiences/all_samples.jsonl"),
help="JSONL 路径(相对仓库根或绝对路径)",
)
parser.add_argument(
"--persist-dir",
default=os.getenv("FEWSHOT_CHROMA_PATH", "data/embeddings/chroma_fewshot"),
help="Chroma 持久化目录",
)
parser.add_argument(
"--force",
action="store_true",
help="清空已有集合并全量重建",
)
args = parser.parse_args()
samples_path = Path(args.samples)
if not samples_path.is_absolute():
samples_path = _REPO_ROOT / samples_path
if not samples_path.is_file():
logger.error("样本文件不存在: %s", samples_path)
return 1
from dotenv import load_dotenv
load_dotenv(_REPO_ROOT / ".env")
from utils.embedding import get_embedder
from utils.fewshot_chroma_store import FewShotChromaStore
from utils.fewshot_selector import ExperienceSample
rows: list[ExperienceSample] = []
with samples_path.open("r", encoding="utf-8") as f:
for line_no, line in enumerate(f, 1):
line = line.strip()
if not line:
continue
try:
rows.append(ExperienceSample.from_dict(json.loads(line)))
except json.JSONDecodeError as e:
logger.warning("跳过第 %s 行 JSON 错误: %s", line_no, e)
if not rows:
logger.error("未解析到任何样本")
return 1
embed_path = os.getenv("EMBEDDING_MODEL_PATH", "").strip() or None
embedder = get_embedder(embed_path)
persist = Path(args.persist_dir)
if not persist.is_absolute():
persist = _REPO_ROOT / persist
store = FewShotChromaStore(embedder, persist_dir=str(persist), persist_to_disk=True)
n = store.build_from_samples(rows, force_rebuild=args.force)
logger.info("完成: 写入 %s 条 → %s", n, persist)
return 0
if __name__ == "__main__":
raise SystemExit(main())