From 37caa87b49b7c87846bf0cb1a27d3b6c48fb86c8 Mon Sep 17 00:00:00 2001 From: marcuspaico Date: Fri, 14 Aug 2026 17:26:26 -0700 Subject: [PATCH] feat: RAG answers with sutta citations (M1) Co-Authored-By: Claude Fable 5 --- scripts/ask.py | 7 +++++++ src/buddhagpt/rag.py | 25 +++++++++++++++++++++++++ tests/test_rag.py | 10 ++++++++++ 3 files changed, 42 insertions(+) create mode 100644 scripts/ask.py create mode 100644 src/buddhagpt/rag.py create mode 100644 tests/test_rag.py diff --git a/scripts/ask.py b/scripts/ask.py new file mode 100644 index 0000000..c9b8796 --- /dev/null +++ b/scripts/ask.py @@ -0,0 +1,7 @@ +import sys +from pathlib import Path +from buddhagpt.rag import answer + +r = answer(sys.argv[1], "mlx-community/Qwen2.5-7B-Instruct-4bit", Path("data/lancedb")) +print(r["answer"]) +print("\nSources:", ", ".join(c["uid"] for c in r["citations"])) diff --git a/src/buddhagpt/rag.py b/src/buddhagpt/rag.py new file mode 100644 index 0000000..920140c --- /dev/null +++ b/src/buddhagpt/rag.py @@ -0,0 +1,25 @@ +from pathlib import Path +from buddhagpt.index import search + +SYSTEM = ( + "You are a thoughtful guide grounded in the Pali Canon. Answer with warmth and " + "precision. Base doctrinal claims on the provided passages and cite them by uid " + "(e.g. [mn21]). If the passages do not cover the question, say so plainly." +) + +def build_prompt(question: str, passages: list[dict]) -> list[dict]: + ctx = "\n\n".join(f"[{p['uid']}] {p['title']}:\n{p['chunk']}" for p in passages) + return [ + {"role": "system", "content": SYSTEM}, + {"role": "user", "content": f"Passages:\n{ctx}\n\nQuestion: {question}"}, + ] + +def answer(question: str, model_path: str, db_path: Path) -> dict: + from mlx_lm import load, generate + passages = search(db_path, question, k=4) + model, tokenizer = load(model_path) + prompt = tokenizer.apply_chat_template( + build_prompt(question, passages), tokenize=False, add_generation_prompt=True + ) + text = generate(model, tokenizer, prompt=prompt, max_tokens=500) + return {"answer": text, "citations": [{"uid": p["uid"], "title": p["title"]} for p in passages]} diff --git a/tests/test_rag.py b/tests/test_rag.py new file mode 100644 index 0000000..1f56691 --- /dev/null +++ b/tests/test_rag.py @@ -0,0 +1,10 @@ +from buddhagpt.rag import build_prompt + +def test_build_prompt_includes_passages_and_citation_instruction(): + msgs = build_prompt("What causes suffering?", [ + {"uid": "sn56.11", "title": "Setting the Wheel in Motion", "chunk": "Craving leads to suffering."} + ]) + system = msgs[0]["content"] + assert "sn56.11" in msgs[1]["content"] + assert "cite" in system.lower() + assert msgs[1]["content"].endswith("What causes suffering?")