feat: synthetic instruction generation via OpenRouter deepseek-v4-flash
This commit is contained in:
82
scripts/gen_data.py
Normal file
82
scripts/gen_data.py
Normal file
@@ -0,0 +1,82 @@
|
||||
import 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
|
||||
|
||||
# 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 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)
|
||||
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)
|
||||
print(f"calls: {len(sample)} | failed: {fail_count} ({fail_rate:.1%}) | raw pairs parsed: {raw_count}", flush=True)
|
||||
|
||||
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:
|
||||
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)
|
||||
Reference in New Issue
Block a user