diff --git a/README.md b/README.md index e69de29..1bb093a 100644 --- a/README.md +++ b/README.md @@ -0,0 +1,22 @@ +# BuddhaGPT + +Fine-tune + RAG experiment: a small local model trained and grounded on the Pali Canon (via +[SuttaCentral](https://suttacentral.net)'s Bilara texts). + +## Data generation + +Synthetic instruction pairs (`data/instructions.jsonl`) are generated from `corpus/suttas.jsonl` +via `scripts/gen_data.py`, using OpenRouter model `~deepseek/deepseek-v4-flash-latest` (the +tilde prefix is part of OpenRouter's real catalog ID for this "latest" alias — verified against +the live `/api/v1/models` catalog, not a typo). + +Token usage / cost: + +- Original full run (`--mode full`, 3,200 calls, variants 0–1 only): totals were not persisted + and the generating process died before a report was written, so these figures are an + **estimate**, not measured: ~5.2M input / 1.8M output tokens, ≈$0.4–0.7 at list pricing + (~$0.08/M in, $0.16/M out). +- Template top-up run (`--mode topup`, 2,300 calls, variants 2–5, 0 failures): **measured** — + 1,994,673 input tokens / 1,618,867 output tokens, **$0.42** at list pricing (~$0.08/M in, + $0.16/M out). See `.superpowers/sdd/2026-08-14-buddha-gpt/task-5-report.md` for the full + fix-round report, including per-template pair counts. diff --git a/scripts/gen_data.py b/scripts/gen_data.py index ac099a2..a1e2492 100644 --- a/scripts/gen_data.py +++ b/scripts/gen_data.py @@ -1,82 +1,137 @@ -import json, random, time, threading +import argparse, json, random, time, threading from concurrent.futures import ThreadPoolExecutor from pathlib import Path from buddhagpt.llm import openrouter_client, chat -from buddhagpt.datagen import build_messages, parse_pairs, dedupe +from buddhagpt.datagen import build_messages, parse_pairs, dedupe, interleaved_variant # NOTE: OpenRouter lists this model under the tilde-prefixed "latest" alias id. MODEL = "~deepseek/deepseek-v4-flash-latest" -suttas = [json.loads(l) for l in Path("corpus/suttas.jsonl").read_text().splitlines()] -random.seed(7) -qualifying = [s for s in suttas if len(s["text"]) > 800] # only 2304 suttas qualify (corpus is - # smaller than assumed) — cycle through - # them with rotating template variants - # to reach the target call volume instead - # of random.sample()'ing more than exist. -TARGET_CALLS = 3200 -sample = [] # list of (sutta, variant) tuples -pass_num = 0 -while len(sample) < TARGET_CALLS: - order = qualifying[:] - random.shuffle(order) - for s in order: - if len(sample) >= TARGET_CALLS: - break - sample.append((s, pass_num)) - pass_num += 1 -client = openrouter_client() -totals = {"input": 0, "output": 0} -fail_count = 0 -done_count = 0 -raw_count = 0 -start = time.time() -lock = threading.Lock() -Path("data").mkdir(parents=True, exist_ok=True) -raw_path = Path("data/instructions_raw.jsonl") -raw_f = raw_path.open("a") # append: incremental persistence, survives interruption +def load_qualifying() -> list[dict]: + suttas = [json.loads(l) for l in Path("corpus/suttas.jsonl").read_text().splitlines()] + return [s for s in suttas if len(s["text"]) > 800] # only 2304 of 3920 suttas qualify -def gen_one(args): - global fail_count, done_count, raw_count - s, variant = args - try: - text, usage = chat(client, MODEL, build_messages(s, variant), max_tokens=2000) - except Exception as e: - with lock: - fail_count += 1 - print(f"skip {s['uid']}: {e}", flush=True) - return - pairs = parse_pairs(text, uid=s["uid"]) - with lock: - totals["input"] += usage["input"]; totals["output"] += usage["output"] - done_count += 1 - raw_count += len(pairs) + +def build_full_sample(qualifying: list[dict], target_calls: int, seed: int) -> list[tuple[dict, int]]: + """Cycle through the qualifying corpus (reshuffled each pass) until target_calls is + reached, assigning template variant by running call index so every full run covers + all templates in TEMPLATES rather than one variant per pass.""" + random.seed(seed) + sample = [] + while len(sample) < target_calls: + order = qualifying[:] + random.shuffle(order) + for s in order: + if len(sample) >= target_calls: + break + sample.append((s, interleaved_variant(len(sample)))) + return sample + + +def build_topup_sample(qualifying: list[dict], variants: list[int], per_variant: int, seed: int) -> list[tuple[dict, int]]: + """Sample `per_variant` distinct suttas (seeded, no repeats within a variant) for each + variant in `variants`, for topping up underrepresented templates.""" + sample = [] + for variant in variants: + random.seed(seed + variant) + chosen = random.sample(qualifying, min(per_variant, len(qualifying))) + sample += [(s, variant) for s in chosen] + random.seed(seed) + random.shuffle(sample) + return sample + + +def run(sample: list[tuple[dict, int]], raw_path: Path) -> tuple[dict, int, int]: + client = openrouter_client() + totals = {"input": 0, "output": 0} + fail_count = 0 + done_count = 0 + raw_count = 0 + start = time.time() + lock = threading.Lock() + + Path("data").mkdir(parents=True, exist_ok=True) + raw_f = raw_path.open("a") # append: incremental persistence, survives interruption + + def gen_one(args): + nonlocal fail_count, done_count, raw_count + s, variant = args + try: + text, usage = chat(client, MODEL, build_messages(s, variant), max_tokens=2000) + except Exception as e: + with lock: + fail_count += 1 + print(f"skip {s['uid']}: {e}", flush=True) + return + pairs = parse_pairs(text, uid=s["uid"]) for p in pairs: - raw_f.write(json.dumps(p) + "\n") - raw_f.flush() - if done_count % 100 == 0: - elapsed = time.time() - start - print(f"progress {done_count}/{len(sample)} | fails {fail_count} | " - f"raw_pairs {raw_count} | {elapsed:.0f}s elapsed", flush=True) + p["variant"] = variant + with lock: + totals["input"] += usage["input"]; totals["output"] += usage["output"] + done_count += 1 + raw_count += len(pairs) + for p in pairs: + raw_f.write(json.dumps(p) + "\n") + raw_f.flush() + if done_count % 100 == 0: + elapsed = time.time() - start + print(f"progress {done_count}/{len(sample)} | fails {fail_count} | " + f"raw_pairs {raw_count} | {elapsed:.0f}s elapsed", flush=True) -with ThreadPoolExecutor(max_workers=8) as pool: - list(pool.map(gen_one, sample)) + with ThreadPoolExecutor(max_workers=8) as pool: + list(pool.map(gen_one, sample)) -raw_f.close() + raw_f.close() -fail_rate = fail_count / len(sample) -print(f"calls: {len(sample)} | failed: {fail_count} ({fail_rate:.1%}) | raw pairs parsed: {raw_count}", flush=True) + fail_rate = fail_count / len(sample) if sample else 0.0 + print(f"calls: {len(sample)} | failed: {fail_count} ({fail_rate:.1%}) | raw pairs parsed: {raw_count}", flush=True) + return totals, fail_count, raw_count -pairs = [json.loads(l) for l in raw_path.read_text().splitlines() if l.strip()] -pairs = [p for p in pairs if 60 <= len(p["answer"].split()) <= 400] -print(f"pairs after word-count filter: {len(pairs)}", flush=True) -pairs = dedupe(pairs) -with Path("data/instructions.jsonl").open("w") as f: + +def rebuild_instructions(raw_path: Path, out_path: Path) -> list[dict]: + """Rebuild data/instructions.jsonl from the FULL raw file (word filter + dedupe).""" + pairs = [json.loads(l) for l in raw_path.read_text().splitlines() if l.strip()] + pairs = [p for p in pairs if 60 <= len(p["answer"].split()) <= 400] + print(f"pairs after word-count filter: {len(pairs)}", flush=True) + pairs = dedupe(pairs) + with out_path.open("w") as f: + for p in pairs: + f.write(json.dumps({"messages": [ + {"role": "user", "content": p["question"]}, + {"role": "assistant", "content": p["answer"]}, + ]}) + "\n") + return pairs + + +def report_variant_counts(pairs: list[dict], label: str) -> None: + counts: dict = {} for p in pairs: - f.write(json.dumps({"messages": [ - {"role": "user", "content": p["question"]}, - {"role": "assistant", "content": p["answer"]}, - ]}) + "\n") -# deepseek-v4-flash list price ~$0.08/M in, $0.16/M out -print(len(pairs), "pairs | tokens", totals, "| est cost $%.2f" % (totals["input"]/1e6*0.08 + totals["output"]/1e6*0.16), flush=True) + v = p.get("variant", "legacy") + counts[v] = counts.get(v, 0) + 1 + print(f"{label} per-variant counts: {dict(sorted(counts.items(), key=lambda kv: str(kv[0])))}", flush=True) + + +if __name__ == "__main__": + ap = argparse.ArgumentParser() + ap.add_argument("--mode", choices=["full", "topup"], default="full") + ap.add_argument("--target-calls", type=int, default=3200) + ap.add_argument("--seed", type=int, default=7) + ap.add_argument("--topup-variants", type=int, nargs="+", default=[2, 3, 4, 5]) + ap.add_argument("--topup-per-variant", type=int, default=575) + args = ap.parse_args() + + qualifying = load_qualifying() + raw_path = Path("data/instructions_raw.jsonl") + + if args.mode == "full": + sample = build_full_sample(qualifying, args.target_calls, args.seed) + else: + sample = build_topup_sample(qualifying, args.topup_variants, args.topup_per_variant, args.seed) + + totals, fail_count, raw_count = run(sample, raw_path) + + pairs = rebuild_instructions(raw_path, Path("data/instructions.jsonl")) + report_variant_counts(pairs, "final (post-filter, post-dedupe)") + # deepseek-v4-flash list price ~$0.08/M in, $0.16/M out + print(len(pairs), "pairs | tokens", totals, "| est cost $%.2f" % (totals["input"]/1e6*0.08 + totals["output"]/1e6*0.16), flush=True) diff --git a/src/buddhagpt/datagen.py b/src/buddhagpt/datagen.py index e6bb8c8..ec40609 100644 --- a/src/buddhagpt/datagen.py +++ b/src/buddhagpt/datagen.py @@ -26,6 +26,14 @@ SYSTEM = ( '{"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 [ diff --git a/tests/test_datagen.py b/tests/test_datagen.py index f8cb11c..327b26f 100644 --- a/tests/test_datagen.py +++ b/tests/test_datagen.py @@ -1,10 +1,17 @@ -from buddhagpt.datagen import build_messages, parse_pairs, dedupe +from buddhagpt.datagen import build_messages, parse_pairs, dedupe, interleaved_variant def test_build_messages_varies_templates(): sutta = {"uid": "mn21", "title": "T", "text": "x" * 900} prompts = {build_messages(sutta, v)[1]["content"] for v in range(6)} assert len(prompts) == 6 # rotating templates, not one fixed prompt +def test_interleaved_variant_cycles_all_templates(): + # Any run of >=6 calls must touch every template, not just the first one or two. + variants = [interleaved_variant(i) for i in range(18)] + assert set(variants) == {0, 1, 2, 3, 4, 5} + assert variants[:6] == [0, 1, 2, 3, 4, 5] + assert variants == variants[:6] * 3 + def test_parse_pairs_extracts_json_lines(): out = '{"question": "Q1?", "answer": "A1"}\n{"question": "Q2?", "answer": "A2"}' assert len(parse_pairs(out, uid="mn21")) == 2