fix: complete template coverage in training data + disclose model alias
This commit is contained in:
22
README.md
22
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.
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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 [
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
Reference in New Issue
Block a user