tiny-agent-112m / code /scripts /gen_synth.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
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()