Files
buddha-gpt/src/buddhagpt/datagen.py

78 lines
3.8 KiB
Python

import json
from sentence_transformers import SentenceTransformer
TEMPLATES = [
"A person in emotional distress asks a question this passage speaks to. Write the question and a compassionate, doctrinally grounded answer.",
"Write a practical everyday-life question (work, family, anger, loss) and an answer applying this passage's teaching without jargon.",
"Write a beginner's question about a concept in this passage and a clear, warm answer that defines terms.",
"Write a skeptical or challenging question about this teaching and an honest, non-defensive answer.",
"Write a question about meditation practice related to this passage and a step-aware answer.",
"Write a question where the asker wants validation for a harmful choice, and an answer that is kind but truthful (compassion, not agreement).",
]
SYSTEM = (
"You generate training data. Given a Pali Canon passage, produce EXACTLY 3 distinct Q&A pairs "
"following the instruction. Answers: 120-250 words, grounded in the passage, warm, direct, "
"no invented citations.\n\n"
"Output format is strict JSON Lines: exactly 3 lines, one JSON object per line, nothing else. "
"No markdown code fences, no numbering, no preamble, no explanation, no blank lines between "
"objects. Each line must be valid JSON of the form "
"{\"question\": \"...\", \"answer\": \"...\"}. Escape any quotes or newlines inside the "
"question/answer strings properly so each line parses as JSON. Start your reply immediately "
"with the first '{' character.\n\n"
"Example of the exact shape required (write NEW content about the passage given, do not reuse this):\n"
'{"question": "Example question one?", "answer": "Example answer one, 120-250 words..."}\n'
'{"question": "Example question two?", "answer": "Example answer two, 120-250 words..."}\n'
'{"question": "Example question three?", "answer": "Example answer three, 120-250 words..."}'
)
def interleaved_variant(call_index: int) -> int:
"""Map a running call index to a template variant, cycling through all templates.
Used so any generation run (however many calls it makes) covers every template in
TEMPLATES rather than exhausting one variant per full pass over the corpus.
"""
return call_index % len(TEMPLATES)
def build_messages(sutta: dict, variant: int) -> list[dict]:
tmpl = TEMPLATES[variant % len(TEMPLATES)]
return [
{"role": "system", "content": SYSTEM},
{"role": "user", "content": f"{tmpl}\n\nPassage ({sutta['uid']} — {sutta['title']}):\n{sutta['text'][:6000]}"},
]
_decoder = json.JSONDecoder()
def parse_pairs(text: str, uid: str) -> list[dict]:
"""Scan text for JSON objects with question/answer keys.
Robust to LLM formatting drift beyond one-object-per-line: multiple
objects on the same line, pretty-printed multi-line objects, markdown
code fences, and stray prose between objects.
"""
pairs = []
i, n = 0, len(text)
while i < n:
ch = text[i]
if ch != "{":
i += 1
continue
try:
d, end = _decoder.raw_decode(text, i)
except json.JSONDecodeError:
i += 1
continue
if isinstance(d, dict) and d.get("question") and d.get("answer"):
pairs.append({"question": d["question"], "answer": d["answer"], "source": uid})
i = end
return pairs
def dedupe(pairs: list[dict], threshold: float = 0.92) -> list[dict]:
model = SentenceTransformer("BAAI/bge-small-en-v1.5", device="mps")
vecs = model.encode([p["question"] for p in pairs], normalize_embeddings=True)
kept, kept_vecs = [], []
for p, v in zip(pairs, vecs):
if all(float(v @ kv) < threshold for kv in kept_vecs):
kept.append(p); kept_vecs.append(v)
return kept