fix: complete template coverage in training data + disclose model alias

This commit is contained in:
marcuspaico
2026-08-15 15:50:27 -07:00
parent cc86a97d10
commit d55d9a648c
4 changed files with 160 additions and 68 deletions

View File

@@ -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.

View File

@@ -1,44 +1,61 @@
import json, random, time, threading import argparse, json, random, time, threading
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
from pathlib import Path from pathlib import Path
from buddhagpt.llm import openrouter_client, chat 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. # NOTE: OpenRouter lists this model under the tilde-prefixed "latest" alias id.
MODEL = "~deepseek/deepseek-v4-flash-latest" 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 def load_qualifying() -> list[dict]:
# smaller than assumed) — cycle through suttas = [json.loads(l) for l in Path("corpus/suttas.jsonl").read_text().splitlines()]
# them with rotating template variants return [s for s in suttas if len(s["text"]) > 800] # only 2304 of 3920 suttas qualify
# to reach the target call volume instead
# of random.sample()'ing more than exist.
TARGET_CALLS = 3200 def build_full_sample(qualifying: list[dict], target_calls: int, seed: int) -> list[tuple[dict, int]]:
sample = [] # list of (sutta, variant) tuples """Cycle through the qualifying corpus (reshuffled each pass) until target_calls is
pass_num = 0 reached, assigning template variant by running call index so every full run covers
while len(sample) < TARGET_CALLS: all templates in TEMPLATES rather than one variant per pass."""
random.seed(seed)
sample = []
while len(sample) < target_calls:
order = qualifying[:] order = qualifying[:]
random.shuffle(order) random.shuffle(order)
for s in order: for s in order:
if len(sample) >= TARGET_CALLS: if len(sample) >= target_calls:
break break
sample.append((s, pass_num)) sample.append((s, interleaved_variant(len(sample))))
pass_num += 1 return sample
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) def build_topup_sample(qualifying: list[dict], variants: list[int], per_variant: int, seed: int) -> list[tuple[dict, int]]:
raw_path = Path("data/instructions_raw.jsonl") """Sample `per_variant` distinct suttas (seeded, no repeats within a variant) for each
raw_f = raw_path.open("a") # append: incremental persistence, survives interruption 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 gen_one(args):
global fail_count, done_count, raw_count 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 s, variant = args
try: try:
text, usage = chat(client, MODEL, build_messages(s, variant), max_tokens=2000) text, usage = chat(client, MODEL, build_messages(s, variant), max_tokens=2000)
@@ -48,6 +65,8 @@ def gen_one(args):
print(f"skip {s['uid']}: {e}", flush=True) print(f"skip {s['uid']}: {e}", flush=True)
return return
pairs = parse_pairs(text, uid=s["uid"]) pairs = parse_pairs(text, uid=s["uid"])
for p in pairs:
p["variant"] = variant
with lock: with lock:
totals["input"] += usage["input"]; totals["output"] += usage["output"] totals["input"] += usage["input"]; totals["output"] += usage["output"]
done_count += 1 done_count += 1
@@ -60,23 +79,59 @@ def gen_one(args):
print(f"progress {done_count}/{len(sample)} | fails {fail_count} | " print(f"progress {done_count}/{len(sample)} | fails {fail_count} | "
f"raw_pairs {raw_count} | {elapsed:.0f}s elapsed", flush=True) f"raw_pairs {raw_count} | {elapsed:.0f}s elapsed", flush=True)
with ThreadPoolExecutor(max_workers=8) as pool: with ThreadPoolExecutor(max_workers=8) as pool:
list(pool.map(gen_one, sample)) list(pool.map(gen_one, sample))
raw_f.close() raw_f.close()
fail_rate = fail_count / len(sample) 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) 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] def rebuild_instructions(raw_path: Path, out_path: Path) -> list[dict]:
print(f"pairs after word-count filter: {len(pairs)}", flush=True) """Rebuild data/instructions.jsonl from the FULL raw file (word filter + dedupe)."""
pairs = dedupe(pairs) pairs = [json.loads(l) for l in raw_path.read_text().splitlines() if l.strip()]
with Path("data/instructions.jsonl").open("w") as f: 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: for p in pairs:
f.write(json.dumps({"messages": [ f.write(json.dumps({"messages": [
{"role": "user", "content": p["question"]}, {"role": "user", "content": p["question"]},
{"role": "assistant", "content": p["answer"]}, {"role": "assistant", "content": p["answer"]},
]}) + "\n") ]}) + "\n")
# deepseek-v4-flash list price ~$0.08/M in, $0.16/M out return pairs
print(len(pairs), "pairs | tokens", totals, "| est cost $%.2f" % (totals["input"]/1e6*0.08 + totals["output"]/1e6*0.16), flush=True)
def report_variant_counts(pairs: list[dict], label: str) -> None:
counts: dict = {}
for p in pairs:
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)

View File

@@ -26,6 +26,14 @@ SYSTEM = (
'{"question": "Example question three?", "answer": "Example answer three, 120-250 words..."}' '{"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]: def build_messages(sutta: dict, variant: int) -> list[dict]:
tmpl = TEMPLATES[variant % len(TEMPLATES)] tmpl = TEMPLATES[variant % len(TEMPLATES)]
return [ return [

View File

@@ -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(): def test_build_messages_varies_templates():
sutta = {"uid": "mn21", "title": "T", "text": "x" * 900} sutta = {"uid": "mn21", "title": "T", "text": "x" * 900}
prompts = {build_messages(sutta, v)[1]["content"] for v in range(6)} prompts = {build_messages(sutta, v)[1]["content"] for v in range(6)}
assert len(prompts) == 6 # rotating templates, not one fixed prompt 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(): def test_parse_pairs_extracts_json_lines():
out = '{"question": "Q1?", "answer": "A1"}\n{"question": "Q2?", "answer": "A2"}' out = '{"question": "Q1?", "answer": "A1"}\n{"question": "Q2?", "answer": "A2"}'
assert len(parse_pairs(out, uid="mn21")) == 2 assert len(parse_pairs(out, uid="mn21")) == 2