patient-edu-qa / scripts /gen_router_data.py
chenhaodev's picture
Initial upload: patient-edu-qa harness (code, data, LoRA, router GGUF)
784ea73 verified
Raw History Blame Contribute Delete
2.63 kB
#!/usr/bin/env python3
"""Build router (1-3B) weak-label SFT dataset.
Input : patient utterance (real + coverage-augmented)
Output: structured plan JSON per our router schema:
{"mode","red_flags":[{"type","trigger","severity"}],
"sub_intents":[{"intent","category","level","risk"}]}
Labels come from the rule engine (mode/red_flags) + category taxonomy +
intent taxonomy. This is weak supervision; the model learns the mapping,
with teacher distillation refining ambiguity later.
Output : output/rag_router/ (or data/router/)
"""
import json
import re
import os
import sys
sys.path.insert(0, os.path.join(os.path.dirname(__file__), ".."))
from scripts.red_flag_rules import scan
OUT = "/workspace/TASK18/data/router"
os.makedirs(OUT, exist_ok=True)
def plan_for(utterance):
red, subs, mode = scan(utterance)
plan = {
"mode": mode,
"red_flags": [],
"sub_intents": subs,
}
if mode == "alert":
type_set = set(("chronicity", "acuteness", "progression"))
for f in red:
if f["type"] in type_set or f["type"] != "population":
plan["red_flags"].append(f)
return plan
def format_example(utterance, plan):
return {
"messages": [
{"role": "user",
"content": f"请把下面这句话拆成结构化的患者意图计划。只输出 JSON。\n患者:{utterance}"},
{"role": "assistant", "content": json.dumps(plan, ensure_ascii=False)},
],
"utterance": utterance,
}
def main():
src = "/workspace/TASK18/data/patient_questions_clean.jsonl"
rows = [json.loads(l) for l in open(src, encoding="utf-8")]
records = []
for r in rows:
u = r["question"]
plan = plan_for(u)
records.append(format_example(u, plan))
# labeled enrichment: override category from ground-truth data where already known
with open(os.path.join(OUT, "router_train.jsonl"), "w", encoding="utf-8") as f:
for rec in records:
f.write(json.dumps(rec, ensure_ascii=False) + "\n")
print(f"router records: {len(records)}")
# stats
from collections import Counter
modes = Counter(p["mode"] for p in (json.loads(r["messages"][1]["content"]) for r in records))
intents = Counter()
for r in records:
p = json.loads(r["messages"][1]["content"])
for s in p["sub_intents"]:
intents[s["intent"]] += 1
print("mode:", dict(modes))
print("top intents:", dict(intents.most_common(12)))
print("saved ->", os.path.join(OUT, "router_train.jsonl"))
if __name__ == "__main__":
main()