Files
buddha-gpt/scripts/gen_data.py

138 lines
5.6 KiB
Python

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, interleaved_variant
# NOTE: OpenRouter lists this model under the tilde-prefixed "latest" alias id.
MODEL = "~deepseek/deepseek-v4-flash-latest"
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 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:
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))
raw_f.close()
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
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:
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)