| """MTRAG benchmark (govt domain), with and without retrieval. Metrics: |
| retrieval recall@k, decline rate on answerable vs unanswerable items, and |
| ROUGE-L against the reference answers. Expects a clone of the IBM |
| mt-rag-benchmark repo at $MTRAG_REPO (see example.env). |
| |
| Usage: python run_mtrag.py --mode rag --k 5 (--inspect to check schemas) |
| """ |
| import argparse, glob, json, os |
| from bench_common import (Retriever, generate, context_block, norm_ws, |
| save_outputs, SYSTEM_PREAMBLE, NORAG_PREAMBLE) |
|
|
| REPO = os.environ.get("MTRAG_REPO", "mt-rag-benchmark") |
| SCRATCH = os.environ.get("SCRATCH", ".") |
| DOMAIN = "govt" |
| DECLINE_MARKERS = ["do not answer", "does not answer", "cannot answer", |
| "don't know", "do not know", "not contain", "unable to", |
| "no information", "i'm sorry", "cannot find"] |
|
|
|
|
| def read_jsonl(path): |
| with open(path, encoding="utf-8") as fh: |
| return [json.loads(line) for line in fh if line.strip()] |
|
|
|
|
| def find_task_files(): |
| |
| |
| |
| |
| path = f"{REPO}/mtrag-human/generation_tasks/reference.jsonl" |
| if not os.path.exists(path): |
| raise FileNotFoundError(f"expected task file missing: {path}") |
| return [path] |
|
|
|
|
| def load_corpus(): |
| path = f"{REPO}/corpora/passage_level/{DOMAIN}.jsonl" |
| rows = read_jsonl(path) |
| passages = [] |
| for r in rows: |
| pid = r.get("_id") or r.get("id") or r.get("document_id") |
| text = r.get("text") or r.get("passage") or "" |
| title = r.get("title") or "" |
| assert pid is not None and text, f"unrecognized corpus schema: {list(r)}" |
| passages.append({"text": f"{title}\n{text}".strip(), "id": str(pid)}) |
| print(f"{len(passages)} {DOMAIN} passages") |
| return passages |
|
|
|
|
| def extract_tasks(files): |
| """Normalize task records: conversation turns, final question, reference |
| answer, answerability label, gold passage ids, domain filter.""" |
| tasks = [] |
| for f in files: |
| for r in read_jsonl(f): |
| |
| |
| if DOMAIN not in str(r.get("Collection", "")).lower(): |
| continue |
| turns = (r.get("input") or r.get("conversation") |
| or r.get("messages") or r.get("turns") or []) |
| targets = r.get("targets") or r.get("reference_answers") or [] |
| if isinstance(targets, list) and targets and isinstance(targets[0], dict): |
| ref = targets[0].get("text", "") |
| elif isinstance(targets, list) and targets: |
| ref = str(targets[0]) |
| else: |
| ref = str(targets) if targets else "" |
| answerability = (r.get("Answerability") or r.get("answerability") |
| or r.get("answerability_label") or "") |
| if isinstance(answerability, list): |
| answerability = answerability[0] if answerability else "" |
| gold_ids = [str(c.get("document_id") or c.get("_id") or "") |
| for c in (r.get("contexts") or []) if isinstance(c, dict)] |
| if not turns: |
| continue |
| last_user = None |
| history = [] |
| for t in turns: |
| role = t.get("speaker") or t.get("role") or "" |
| text = t.get("text") or t.get("content") or "" |
| if role.lower() in ("user", "human"): |
| last_user = text |
| history.append(f"{role}: {text}") |
| if not last_user: |
| continue |
| tasks.append({"task_id": r.get("task_id") or r.get("id") or "", |
| "history": history[:-1], "question": last_user, |
| "reference": ref, |
| "answerability": str(answerability).upper(), |
| "gold_ids": gold_ids}) |
| print(f"{len(tasks)} {DOMAIN} tasks from {len(files)} files") |
| return tasks |
|
|
|
|
| def rouge_l_f(a, b): |
| """Plain LCS-based ROUGE-L F1, no external deps.""" |
| ta, tb = norm_ws(a).split(), norm_ws(b).split() |
| if not ta or not tb: |
| return 0.0 |
| dp = [[0] * (len(tb) + 1) for _ in range(len(ta) + 1)] |
| for i in range(len(ta)): |
| for j in range(len(tb)): |
| dp[i + 1][j + 1] = (dp[i][j] + 1 if ta[i] == tb[j] |
| else max(dp[i][j + 1], dp[i + 1][j])) |
| lcs = dp[-1][-1] |
| p, r = lcs / len(tb), lcs / len(ta) |
| return 2 * p * r / (p + r) if p + r else 0.0 |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--mode", choices=["rag", "norag"]) |
| ap.add_argument("--k", type=int, default=5) |
| ap.add_argument("--inspect", action="store_true") |
| ap.add_argument("--limit", type=int, default=0) |
| args = ap.parse_args() |
|
|
| files = find_task_files() |
| if args.inspect: |
| print("task files found:") |
| for f in files: |
| print(" ", f) |
| rows = read_jsonl(f) |
| if rows: |
| print(" first record keys:", sorted(rows[0].keys())) |
| return |
| assert args.mode, "--mode required unless --inspect" |
|
|
| tasks = extract_tasks(files) |
| if args.limit: |
| tasks = tasks[:args.limit] |
|
|
| retriever = None |
| if args.mode == "rag": |
| retriever = Retriever(load_corpus(), f"{SCRATCH}/emb/mtrag_{DOMAIN}.npy") |
|
|
| samples = [] |
| for t in tasks: |
| hist = "\n".join(t["history"][-6:]) |
| instruction = (f"Conversation so far:\n{hist}\n\n" |
| f"Current question: {t['question']}") |
| if args.mode == "rag": |
| hits = retriever.search(t["question"], k=args.k) |
| block = f"Context:\n{context_block(hits)}\n\n{instruction}" |
| resp = generate(block, SYSTEM_PREAMBLE) |
| retrieved_ids = [h["id"] for h, _ in hits] |
| recall = (len(set(retrieved_ids) & set(t["gold_ids"])) |
| / len(set(t["gold_ids"]))) if t["gold_ids"] else None |
| else: |
| resp = generate(instruction, NORAG_PREAMBLE) |
| retrieved_ids, recall = [], None |
| declined = any(m in resp.lower() for m in DECLINE_MARKERS) |
| unanswerable = t["answerability"] in ("UNANSWERABLE", "NO") |
| rouge = rouge_l_f(resp, t["reference"]) if t["reference"] else None |
| samples.append({**{k: t[k] for k in |
| ("task_id", "question", "reference", "answerability")}, |
| "response": resp, "declined": declined, |
| "unanswerable": unanswerable, |
| "retrieved": retrieved_ids, "recall_at_k": recall, |
| "rouge_l": rouge}) |
|
|
| unans = [s for s in samples if s["unanswerable"]] |
| ans = [s for s in samples if not s["unanswerable"]] |
| recs = [s["recall_at_k"] for s in samples if s["recall_at_k"] is not None] |
| rls = [s["rouge_l"] for s in ans if s["rouge_l"] is not None] |
| metrics = { |
| "n_answerable": len(ans), "n_unanswerable": len(unans), |
| "decline_rate_on_unanswerable": |
| (sum(s["declined"] for s in unans) / len(unans)) if unans else None, |
| "decline_rate_on_answerable": |
| (sum(s["declined"] for s in ans) / len(ans)) if ans else None, |
| "rouge_l_answerable": (sum(rls) / len(rls)) if rls else None, |
| "mean_recall_at_k": (sum(recs) / len(recs)) if recs else None, |
| } |
| save_outputs(f"{SCRATCH}/results", "mtrag", args.mode, metrics, samples, |
| extra={"domain": DOMAIN, "k": args.k}) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|