"""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 (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 , 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()