Download code/scripts/gen_synth.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 10.8 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/gen_synth.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/gen_synth.py
-
curl -L -o gen_synth.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/gen_synth.py
10.8 kB
| """Build the synthetic sources as tokenized .bin files (same layout as tokenize_data.py). | |
| synth_tools scripted-solver trajectories on fresh invented-fact workspaces (ta-v1 format) | |
| synth_grounded reading comprehension: answer is a span copied from the given text, or NOT_FOUND | |
| (SQuAD v2, HotpotQA distractor, and workspace files rendered as passages) | |
| synth_reasoning math word problems with the solution inside <think> (OpenMathInstruct-2, GSM8K) | |
| Task seeds >= 1_000_000 are used here; eval/RL tasks use seeds below that, so they never overlap. | |
| Usage: source env.sh && $TA_PY scripts/gen_synth.py --tools 300000 | |
| """ | |
| import argparse | |
| import glob | |
| import json | |
| import os | |
| import random | |
| import re | |
| from multiprocessing import Pool | |
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") | |
| import numpy as np | |
| import pyarrow.parquet as pq | |
| from tokenizers import Tokenizer | |
| from tiny_agent.chat import render, repeated_calls | |
| from tiny_agent.tasks import EXTRA_TRAIN_KINDS, TRAIN_KINDS, make_task, oracle_trajectory, vary_question | |
| from tiny_agent.text import DATA, EOS_ID | |
| TRAIN_SEED0 = 1_000_000 | |
| VAL_EVERY = 200 | |
| READ_SYS = "Answer using only the given text. Copy the answer span exactly. If the text does not contain the answer, answer NOT_FOUND." | |
| MATH_SYS = "Solve the problem. Reason step by step inside <think>, then give the final answer." | |
| _tok = None | |
| def tok(): | |
| global _tok | |
| if _tok is None: | |
| _tok = Tokenizer.from_file(f"{DATA}/tokenizer.json") | |
| return _tok | |
| class BinWriter: | |
| def __init__(self, source, part): | |
| self.paths = {s: f"{DATA}/tok/{s}/{source}/{part}.bin" for s in ("train", "val")} | |
| for p in self.paths.values(): | |
| os.makedirs(os.path.dirname(p), exist_ok=True) | |
| self.f = {s: open(p + ".tmp", "wb") for s, p in self.paths.items()} | |
| self.k, self.n = 0, {"train": 0, "val": 0} | |
| def add_many(self, texts): | |
| for enc in tok().encode_batch(texts, add_special_tokens=False): | |
| arr = np.asarray(enc.ids + [EOS_ID], dtype=np.uint16) | |
| split = "val" if self.k % VAL_EVERY == VAL_EVERY - 1 else "train" | |
| arr.tofile(self.f[split]) | |
| self.n[split] += arr.size | |
| self.k += 1 | |
| def close(self): | |
| for s, f in self.f.items(): | |
| f.close() | |
| os.replace(self.paths[s] + ".tmp", self.paths[s]) | |
| return self.n | |
| def tools_part(args): | |
| part, start, n, name = args | |
| w = BinWriter(name, f"part{part:03d}") | |
| buf, bad = [], 0 | |
| for i in range(start, start + n): | |
| rng = random.Random(TRAIN_SEED0 + i) | |
| task = make_task(rng, TRAIN_KINDS[i % len(TRAIN_KINDS)]) | |
| msgs, ok = oracle_trajectory(task, rng) | |
| if not ok: | |
| bad += 1 | |
| continue | |
| buf.append(render(msgs)) | |
| if len(buf) == 256: | |
| w.add_many(buf) | |
| buf = [] | |
| if buf: | |
| w.add_many(buf) | |
| return name, part, w.close(), bad | |
| # warm-start data for r2: the new training kinds (save_value = the only `write` demos, service_hop) | |
| # upweighted to ~1/3, every question reworded by vary_question. Seeds 3,000,000+ (scripted range, | |
| # clear of synth_tools' 1,000,000-1,299,999 and the teacher's 5,000,000+). | |
| AGENT_V2_KINDS = TRAIN_KINDS + EXTRA_TRAIN_KINDS * 2 | |
| AGENT_V2_SEED0 = 3_000_000 | |
| def agent_v2_part(args): | |
| part, start, n, name = args | |
| w = BinWriter(name, f"part{part:03d}") | |
| buf, bad = [], 0 | |
| for i in range(start, start + n): | |
| rng = random.Random(AGENT_V2_SEED0 + i) | |
| task = vary_question(make_task(rng, AGENT_V2_KINDS[i % len(AGENT_V2_KINDS)]), rng) | |
| msgs, ok = oracle_trajectory(task, rng) | |
| if not ok: | |
| bad += 1 | |
| continue | |
| buf.append(render(msgs)) | |
| if len(buf) == 256: | |
| w.add_many(buf) | |
| buf = [] | |
| if buf: | |
| w.add_many(buf) | |
| return name, part, w.close(), bad | |
| def teacher_v2_texts(pattern, max_tokens=3800): | |
| """Kept teacher episodes with the question reworded (task regenerated from its seed): Qwen's | |
| natural reasoning, paired with wording the templates never produced.""" | |
| out = [] | |
| for path in sorted(glob.glob(pattern)): | |
| for line in open(path): | |
| r = json.loads(line) | |
| if teacher_reject(r, max_tokens)[0]: | |
| continue | |
| rng = random.Random(r["seed"] * 7 + 1) | |
| t = vary_question(make_task(random.Random(r["seed"]), r["kind"]), rng, p=0.9) | |
| msgs = [dict(m) for m in r["messages"]] | |
| msgs[1] = {"role": "user", "content": t.question} | |
| text = render(msgs) | |
| if len(tok().encode(text, add_special_tokens=False).ids) <= max_tokens: | |
| out.append(text) | |
| random.Random(4).shuffle(out) | |
| print("teacher_v2", len(out), flush=True) | |
| return out | |
| def chat(system, user, think, answer): | |
| return render([{"role": "system", "content": system}, {"role": "user", "content": user}, | |
| {"role": "assistant", "think": think, "content": answer, "tool_calls": []}]) | |
| def grounded_texts(seed): | |
| rng = random.Random(seed) | |
| out = [] | |
| for f in glob.glob(f"{DATA}/raw/extra/squad_v2/**/*.parquet", recursive=True): | |
| for r in pq.read_table(f).to_pylist(): | |
| ans = r["answers"]["text"] | |
| out.append(chat(READ_SYS, f"{r['context']}\n\nQuestion: {r['question']}", None, ans[0] if ans else "NOT_FOUND")) | |
| for f in glob.glob(f"{DATA}/raw/extra/hotpot_qa/**/*.parquet", recursive=True): | |
| for r in pq.read_table(f).to_pylist(): | |
| if r["answer"].lower() in ("yes", "no"): | |
| continue | |
| ctx = "\n\n".join(f"{t}: {''.join(s)}" for t, s in zip(r["context"]["title"], r["context"]["sentences"])) | |
| sup = sorted(set(r["supporting_facts"]["title"])) | |
| out.append(chat(READ_SYS, f"{ctx}\n\nQuestion: {r['question']}", | |
| f"The relevant paragraphs are {' and '.join(sup)}.", r["answer"])) | |
| # workspace files as passages (invented facts, so only the text can answer) | |
| for i in range(60000): | |
| r2 = random.Random(TRAIN_SEED0 * 2 + i) | |
| task = make_task(r2, r2.choice(["config_value", "code_constant", "csv_lookup", "doc_fact", "not_found"])) | |
| f = task.meta.get("file") or r2.choice(sorted(task.files)) | |
| if f not in task.files: | |
| f = r2.choice(sorted(task.files)) | |
| out.append(chat(READ_SYS, f"File: {f}\n{task.files[f]}\nQuestion: {task.question}", None, | |
| task.answer if task.meta.get("file") == f else "NOT_FOUND")) | |
| rng.shuffle(out) | |
| return out | |
| def math_texts(seed): | |
| out = [] | |
| for f in glob.glob(f"{DATA}/raw/extra/gsm8k/**/*.parquet", recursive=True): | |
| for r in pq.read_table(f).to_pylist(): | |
| sol, ans = r["answer"].split("####") | |
| sol = re.sub(r"<<[^>]*>>", "", sol).strip() | |
| out.append(chat(MATH_SYS, r["question"], sol, f"The answer is {ans.strip()}.")) | |
| for f in sorted(glob.glob(f"{DATA}/raw/extra/OpenMathInstruct-2/**/*.parquet", recursive=True)): | |
| t = pq.read_table(f, columns=["problem", "generated_solution", "expected_answer"]).to_pylist() | |
| for r in t: | |
| sol = r["generated_solution"] | |
| if len(sol) > 3000: | |
| continue | |
| out.append(chat(MATH_SYS, r["problem"], sol, f"The answer is {r['expected_answer']}.")) | |
| random.Random(seed).shuffle(out) | |
| return out | |
| def write_texts(source, texts, parts=8): | |
| jobs = [(source, i, texts[i::parts]) for i in range(parts)] | |
| with Pool(parts) as pool: | |
| for r in pool.imap_unordered(_write_job, jobs): | |
| print(r, flush=True) | |
| def _write_job(args): | |
| source, i, texts = args | |
| w = BinWriter(source, f"part{i:03d}") | |
| for j in range(0, len(texts), 512): | |
| w.add_many(texts[j:j + 512]) | |
| return source, i, w.close() | |
| def teacher_reject(r, max_tokens=3800): | |
| """Why a teacher episode is unusable (None = keep): incorrect/ungrounded, malformed or repeated | |
| calls, or too long for the student's context. Returns (reason, rendered text).""" | |
| if not (r["ok"] and r["grounded"]): | |
| return "wrong", None | |
| if any("error" in c for m in r["messages"] if m["role"] == "assistant" for c in m.get("tool_calls") or []): | |
| return "bad", None | |
| if sum(repeated_calls(r["messages"])): | |
| return "repeat", None | |
| text = render(r["messages"]) | |
| if len(tok().encode(text, add_special_tokens=False).ids) > max_tokens: | |
| return "long", text | |
| return None, text | |
| def teacher_texts(pattern, max_tokens=3800): | |
| """Kept teacher episodes: correct, grounded, no malformed or repeated calls, fits the student's context.""" | |
| out, stats = [], {"total": 0, "kept": 0, "wrong": 0, "bad": 0, "repeat": 0, "long": 0} | |
| for path in sorted(glob.glob(pattern)): | |
| for line in open(path): | |
| r = json.loads(line) | |
| stats["total"] += 1 | |
| why, text = teacher_reject(r, max_tokens) | |
| if why: | |
| stats[why] += 1 | |
| continue | |
| out.append(text) | |
| stats["kept"] += 1 | |
| print("teacher", stats, flush=True) | |
| random.Random(3).shuffle(out) | |
| return out | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--tools", type=int, default=300000) | |
| ap.add_argument("--agent_v2", type=int, default=48000) | |
| ap.add_argument("--workers", type=int, default=24) | |
| ap.add_argument("--only", default="tools,grounded,reasoning") | |
| ap.add_argument("--tools_name", default="synth_tools", help="output source dir (write elsewhere, then swap)") | |
| ap.add_argument("--teacher_glob", default=f"{DATA}/teacher/*.jsonl") | |
| a = ap.parse_args() | |
| only = a.only.split(",") | |
| if "tools" in only: | |
| per = a.tools // a.workers | |
| with Pool(a.workers) as pool: | |
| for r in pool.imap_unordered(tools_part, [(i, i * per, per, a.tools_name) for i in range(a.workers)]): | |
| print(r, flush=True) | |
| if "grounded" in only: | |
| write_texts("synth_grounded", grounded_texts(1)) | |
| if "reasoning" in only: | |
| write_texts("synth_reasoning", math_texts(2)) | |
| if "agent_v2" in only: | |
| per = a.agent_v2 // a.workers | |
| with Pool(a.workers) as pool: | |
| for r in pool.imap_unordered(agent_v2_part, [(i, i * per, per, "synth_agent_v2") for i in range(a.workers)]): | |
| print(r, flush=True) | |
| if "teacher_v2" in only: | |
| texts = teacher_v2_texts(a.teacher_glob) | |
| if texts: | |
| write_texts("synth_teacher_v2", texts, parts=4) | |
| if "teacher" in only: | |
| texts = teacher_texts(a.teacher_glob) | |
| if texts: | |
| write_texts("synth_teacher", texts, parts=4) | |
| if __name__ == "__main__": | |
| main() | |