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 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)

View File

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

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():
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