patient-edu-qa / scripts /build_multi_rag.py
chenhaodev's picture
Initial upload: patient-edu-qa harness (code, data, LoRA, router GGUF)
784ea73 verified
Raw History Blame Contribute Delete
3.18 kB
#!/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]}")