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,82 +1,137 @@
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
# 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) def load_qualifying() -> list[dict]:
raw_path = Path("data/instructions_raw.jsonl") suttas = [json.loads(l) for l in Path("corpus/suttas.jsonl").read_text().splitlines()]
raw_f = raw_path.open("a") # append: incremental persistence, survives interruption 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 def build_full_sample(qualifying: list[dict], target_calls: int, seed: int) -> list[tuple[dict, int]]:
s, variant = args """Cycle through the qualifying corpus (reshuffled each pass) until target_calls is
try: reached, assigning template variant by running call index so every full run covers
text, usage = chat(client, MODEL, build_messages(s, variant), max_tokens=2000) all templates in TEMPLATES rather than one variant per pass."""
except Exception as e: random.seed(seed)
with lock: sample = []
fail_count += 1 while len(sample) < target_calls:
print(f"skip {s['uid']}: {e}", flush=True) order = qualifying[:]
return random.shuffle(order)
pairs = parse_pairs(text, uid=s["uid"]) for s in order:
with lock: if len(sample) >= target_calls:
totals["input"] += usage["input"]; totals["output"] += usage["output"] break
done_count += 1 sample.append((s, interleaved_variant(len(sample))))
raw_count += len(pairs) 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: for p in pairs:
raw_f.write(json.dumps(p) + "\n") p["variant"] = variant
raw_f.flush() with lock:
if done_count % 100 == 0: totals["input"] += usage["input"]; totals["output"] += usage["output"]
elapsed = time.time() - start done_count += 1
print(f"progress {done_count}/{len(sample)} | fails {fail_count} | " raw_count += len(pairs)
f"raw_pairs {raw_count} | {elapsed:.0f}s elapsed", flush=True) 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: 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:
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: for p in pairs:
f.write(json.dumps({"messages": [ v = p.get("variant", "legacy")
{"role": "user", "content": p["question"]}, counts[v] = counts.get(v, 0) + 1
{"role": "assistant", "content": p["answer"]}, print(f"{label} per-variant counts: {dict(sorted(counts.items(), key=lambda kv: str(kv[0])))}", flush=True)
]}) + "\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) 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