#!/usr/bin/env python3 """Build per-category FAISS indexes (multi-RAG) from UpToDate patient Q&A. One index per category (30). Each doc = question + answer (English UpToDate). Query language = Chinese (patients) -> use a multilingual embedding model. """ import json import os import numpy as np os.environ.setdefault("HF_ENDPOINT", "https://hf-mirror.com") ROOT = "/workspace/TASK18" QA = os.path.join(ROOT, "data/patient_education_qa.jsonl") OUT = os.path.join(ROOT, "data/rag") MODEL_NAME = "BAAI/bge-m3" # multilingual (zh query -> en docs) os.makedirs(OUT, exist_ok=True) from sentence_transformers import SentenceTransformer import faiss model = SentenceTransformer(MODEL_NAME) model.max_seq_length = 512 print("model loaded:", MODEL_NAME, "| device:", model.device, flush=True) rows = [json.loads(l) for l in open(QA, encoding="utf-8")] # group by category from collections import defaultdict bycat = defaultdict(list) for r in rows: bycat[r["category"]].append(r) print("categories:", len(bycat), flush=True) per_cat_stats = {} for cat, items in bycat.items(): # dedup by question within category seen = set() uniq = [] for r in items: if r["question"] in seen: continue seen.add(r["question"]) uniq.append(r) docs = [] meta = [] for r in uniq: d = f"{r['question']}\n{r['answer']}" docs.append(d) meta.append({"slug": r["slug"], "id": r["id"], "question": r["question"], "title": r["title"], "level": r["level"], "category": cat}) emb = model.encode(docs, normalize_embeddings=True, show_progress_bar=True, batch_size=64) emb = np.asarray(emb, dtype="float32") idx = faiss.IndexFlatIP(emb.shape[1]) idx.add(emb) sub = os.path.join(OUT, cat) os.makedirs(sub, exist_ok=True) faiss.write_index(idx, os.path.join(sub, "index.faiss")) with open(os.path.join(sub, "docs.json"), "w", encoding="utf-8") as f: json.dump(docs, f, ensure_ascii=False) with open(os.path.join(sub, "meta.json"), "w", encoding="utf-8") as f: json.dump(meta, f, ensure_ascii=False) per_cat_stats[cat] = len(uniq) print(f" [{cat}] docs={len(uniq)}", flush=True) with open(os.path.join(OUT, "index_map.json"), "w", encoding="utf-8") as f: json.dump(per_cat_stats, f, ensure_ascii=False, indent=1) # ---- sanity: multilingual cross-lingual query ---- print("\n=== cross-lingual sanity (Chinese query -> English docs) ===") q = "我最近2周一直头疼,有什么免吃药的方法缓解" qe = model.encode([q], normalize_embeddings=True) for cat in ["brain-and-nerves", "gastrointestinal-system"]: sub = os.path.join(OUT, cat) if not os.path.exists(os.path.join(sub, "index.faiss")): continue idx = faiss.read_index(os.path.join(sub, "index.faiss")) meta = json.load(open(os.path.join(sub, "meta.json"), encoding="utf-8")) docs = json.load(open(os.path.join(sub, "docs.json"), encoding="utf-8")) D, I = idx.search(qe, 3) print(f"\n-- category: {cat} --") for d, i in zip(D[0], I[0]): print(f" score={d:.3f} | {meta[i]['question'][:70]}")