darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
14.8 kB
"""GRPO directly on a base snapshot (no SFT stage, per RL Excursions), on workspace tasks.
Rollouts (tiny_agent.rollout):
continuous (default): slots are refilled as episodes end; a step trains on the first n_tasks
completed groups that carry signal, and episodes still running keep going across the update
(partial rollouts), up to --max_stale policy versions old.
lockstep: one batch of n_tasks x G rollouts per step (fallback).
Rewards and credit (tiny_agent.rl, MiMo-V2.6 reference profile): correct/grounded reward, in-group
length penalty on passes, advantage = r - group mean, groups without signal dropped, segment
penalty for bad tool-call turns with signed rebalance.
Loss: prompt-mean aggregation; per-token importance ratio sg(pi/mu) against the recorded sampling
log-prob mu, with tokens outside [--clip_lo, --clip_hi] masked; no KL.
source env.sh && $TA_PY scripts/grpo.py --ckpt $TA_DATA/runs/main/final.pt --out $TA_DATA/rl/r1 --steps 300
"""
import argparse
import json
import os
import random
import time
from collections import defaultdict, deque
from dataclasses import asdict
import torch
import torch.nn.functional as F
from tokenizers import Tokenizer
from scripts.eval_agent import evaluate
from tiny_agent.checkpoint import load_model
from tiny_agent.model import make_block_mask
from tiny_agent.optim import build_optimizers
from tiny_agent.rl import (LengthPenalty, base_reward, group_advantages, has_signal, length_deltas,
signed_rebalance)
from tiny_agent.rollout import make_roller
from tiny_agent.tasks import TRAIN_KINDS, make_task, normalize, vary_question
from tiny_agent.text import DATA, PAD_ID
RL_SEED0 = 100_000
BUCKETS = (1024, 2048, 3072, 4096)
def shape_group(eps, lp_cfg, invented_penalty=0.0):
"""Sets e.reward (base reward + length penalty); returns the rewards."""
base = [base_reward(e.correct, e.grounded, invented=e.invented, invented_penalty=invented_penalty) for e in eps]
deltas = length_deltas(base, [e.signals() for e in eps], lp_cfg)
for e, b, d in zip(eps, base, deltas):
e.reward = b + d
return [e.reward for e in eps]
def token_weights(groups, kappa, min_scale, max_scale):
"""Per episode: (episode, weight per token position for predicting tokens[1:]).
Prompt-mean: each group's weights are divided by its generated-token count and by the number
of groups, so every prompt counts equally whatever its rollout lengths."""
eps, adv, nf, nc = [], [], [], []
for g in groups:
for e, a in zip(g, group_advantages([e.reward for e in g])):
flagged = sum(b - s for (s, b), f in zip(e.turn_spans, e.flags) if f)
eps.append(e)
adv.append(a)
nf.append(flagged)
nc.append(e.gen_tokens - flagged)
scales = signed_rebalance(adv, nf, nc, kappa, min_scale, max_scale)
gtok = {id(e): sum(x.gen_tokens for x in g) for g in groups for e in g}
out = []
for e, a, (wc, wf) in zip(eps, adv, scales):
if a == 0.0:
continue
norm = 1.0 / (len(groups) * max(1, gtok[id(e)]))
w = [a * wc * norm * m for m in e.gen_mask]
for (s, b), f in zip(e.turn_spans, e.flags):
if f:
for t in range(s, b):
w[t] = a * wf * norm
out.append((e, w[1:]))
return out
def policy_loss(model, items, micro, device, T, clip_lo, clip_hi):
"""-sum(w * sg(pi/mu) * mask * log pi) over generated tokens. `model` is compiled with
dynamic=False: rows are padded to `micro` and lengths to a few buckets so shapes stay static."""
items = sorted(items, key=lambda it: len(it[0].tokens), reverse=True)
st = defaultdict(float)
for i in range(0, len(items), micro):
chunk = items[i:i + micro]
need = min(T, max(len(e.tokens) for e, _ in chunk))
L = next((b for b in BUCKETS if b >= need), T)
x = torch.full((micro, L), PAD_ID, dtype=torch.long)
tgt = torch.full((micro, L), -1, dtype=torch.long)
w = torch.zeros(micro, L)
mu = torch.zeros(micro, L)
g = torch.zeros(micro, L)
doc = torch.ones(micro, L, dtype=torch.long)
for r, (e, wt) in enumerate(chunk):
n = min(len(e.tokens), L)
x[r, :n] = torch.tensor(e.tokens[:n])
doc[r, :n] = 0
tgt[r, : n - 1] = torch.tensor(e.tokens[1:n])
w[r, : n - 1] = torch.tensor(wt[: n - 1])
mu[r, : n - 1] = torch.tensor(e.logp[1:n])
g[r, : n - 1] = torch.tensor(e.gen_mask[1:n], dtype=torch.float32)
x, tgt, w, mu, g, doc = (t.to(device) for t in (x, tgt, w, mu, g, doc))
with torch.autocast("xpu", dtype=torch.bfloat16):
logits = model(x, doc, make_block_mask(doc, getattr(model, "_orig_mod", model).cfg.swa_window))
logp = -F.cross_entropy(logits.reshape(-1, logits.size(-1)).float(), tgt.reshape(-1).clamp_min(0),
reduction="none").view_as(w)
diff = (logp.detach() - mu) * g
ratio = torch.exp(diff.clamp(-20, 20))
keep = ((ratio >= clip_lo) & (ratio <= clip_hi)).float()
loss = -(w * ratio * keep * logp).sum()
loss.backward()
st["loss"] += loss.item()
st["gen"] += g.sum().item()
st["abs_logratio"] += diff.abs().sum().item()
st["clipped"] += (g * (1 - keep)).sum().item()
gen = max(1.0, st["gen"])
return {"loss": round(st["loss"], 6), "mismatch": round(st["abs_logratio"] / gen, 5),
"clip_frac": round(st["clipped"] / gen, 5)}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--ckpt", required=True)
ap.add_argument("--out", required=True)
ap.add_argument("--steps", type=int, default=300)
ap.add_argument("--n_tasks", type=int, default=32, help="groups with signal per update")
ap.add_argument("--G", type=int, default=8)
ap.add_argument("--opt", default="adamw", choices=["adamw", "muon"])
ap.add_argument("--lr", type=float, default=5e-6)
ap.add_argument("--max_len", type=int, default=4096)
ap.add_argument("--max_turn_tokens", type=int, default=768)
ap.add_argument("--micro", type=int, default=4)
ap.add_argument("--engine", default="continuous", choices=["continuous", "lockstep"])
ap.add_argument("--slots", type=int, default=256)
ap.add_argument("--prefill_rows", type=int, default=16)
ap.add_argument("--prefill_wait", type=int, default=8)
ap.add_argument("--max_stale", type=int, default=2, help="drop groups that started this many updates ago")
ap.add_argument("--clip_lo", type=float, default=0.2)
ap.add_argument("--clip_hi", type=float, default=5.0)
ap.add_argument("--kappa", type=float, default=2.0)
ap.add_argument("--length_penalty", type=float, default=0.2, help="max penalty; 0 disables")
ap.add_argument("--vary", type=float, default=0.7, help="probability of rewording a training question (0 = templates only)")
ap.add_argument("--step0", type=int, default=0, help="resume numbering after a restart from rl_step<step0>.pt")
ap.add_argument("--invented_penalty", type=float, default=0.3,
help="reward for a wrong answer seen in no tool result (vs 0 for other wrong answers)")
ap.add_argument("--eval_every", type=int, default=25)
ap.add_argument("--eval_tasks", type=int, default=150)
ap.add_argument("--kinds", default=",".join(TRAIN_KINDS), help="held-out kinds are eval-only by default")
ap.add_argument("--curriculum", type=int, default=1)
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--debug_reward", type=int, default=0, help="random correctness (smoke tests only)")
a = ap.parse_args()
os.makedirs(a.out, exist_ok=True)
dev = "xpu"
tok = Tokenizer.from_file(f"{DATA}/tokenizer.json")
model = load_model(a.ckpt)
params = [p for p in model.parameters() if p.requires_grad]
if a.opt == "muon":
opts = build_optimizers(model, lr=a.lr, weight_decay=0.0)
else:
opts = [torch.optim.AdamW(params, lr=a.lr, betas=(0.9, 0.99), weight_decay=0.0)]
roller = make_roller("lockstep" if a.engine == "lockstep" else "continuous", model, tok, device=dev,
max_len=a.max_len, max_turn_tokens=a.max_turn_tokens, slots=a.slots,
prefill_rows=a.prefill_rows, prefill_wait=a.prefill_wait)
cmodel = torch.compile(model, dynamic=False)
lp_cfg = LengthPenalty(max_penalty=a.length_penalty)
kinds = a.kinds.split(",")
rng = random.Random(a.seed)
# curriculum: per-kind success EMA (updated from every completed group, before filtering);
# sample kinds whose groups are likely to have signal, with a floor so none disappear
succ = {k: 0.5 for k in kinds}
def pick_kind():
if not a.curriculum:
return rng.choice(kinds)
return rng.choices(kinds, weights=[max(succ[k] * (1 - succ[k]), 0.04) for k in kinds])[0]
def new_task():
t = make_task(random.Random(RL_SEED0 + rng.randrange(900_000)), pick_kind())
return vary_question(t, rng, a.vary) if a.vary > 0 else t
log = open(os.path.join(a.out, "log.jsonl"), "a")
def write(rec):
log.write(json.dumps(rec) + "\n")
log.flush()
print(json.dumps(rec), flush=True)
def do_eval(step):
model.eval()
write({"step": step, "eval": evaluate(model, tok, a.eval_tasks, k=4, max_len=a.max_len)})
ready, pending, gid = deque(), defaultdict(list), 0
cnt = defaultdict(int)
def complete(g):
"""A finished group: curriculum update, then keep it only if it carries signal."""
if a.debug_reward:
for e in g:
e.correct = e.grounded = rng.random() < 0.5
k = g[0].task.kind
succ[k] = 0.9 * succ[k] + 0.1 * (sum(e.correct for e in g) / len(g))
cnt["groups"] += 1
if has_signal(shape_group(g, lp_cfg, a.invented_penalty)):
ready.append(g)
else:
cnt["no_signal"] += 1
if a.step0 == 0:
do_eval(0)
for step in range(a.step0 + 1, a.steps + 1):
t0 = time.time()
model.eval()
c0 = dict(cnt)
if a.engine == "lockstep":
tasks = [new_task() for _ in range(a.n_tasks)]
eps = roller.run([t for t in tasks for _ in range(a.G)])
for i in range(a.n_tasks):
complete(eps[i * a.G:(i + 1) * a.G])
else:
warned = 0
while len(ready) < a.n_tasks:
while roller.in_flight() < a.slots * 3 // 2: # keep a queue so freed slots refill at once
roller.submit([new_task()] * a.G, group=gid)
gid += 1
for e in roller.tick():
pending[e.group].append(e)
if len(pending[e.group]) == a.G:
complete(pending.pop(e.group))
done = cnt["groups"] - c0.get("groups", 0)
if done >= (warned + 1) * 10 * a.n_tasks:
warned += 1
print(f"step {step}: {done} groups, only {len(ready)} with signal", flush=True)
if ready or warned >= 3:
break
batch = []
while ready and len(batch) < a.n_tasks:
g = ready.popleft()
if roller.version - min(e.version for e in g) > a.max_stale:
cnt["stale"] += 1
continue
batch.append(g)
t_roll = time.time() - t0
if not batch:
write({"step": step, "skipped": "no group with signal"})
continue
model.train()
items = token_weights(batch, a.kappa, 0.5, 2.0)
st = policy_loss(cmodel, items, a.micro, dev, a.max_len, a.clip_lo, a.clip_hi)
gnorm = torch.nn.utils.clip_grad_norm_(params, 1.0).item()
for o in opts:
o.step()
o.zero_grad(set_to_none=True)
roller.version += 1
flat = [e for g in batch for e in g]
per_kind = defaultdict(list)
for e in flat:
per_kind[e.task.kind].append(e.correct)
n_turns = sum(e.turns for e in flat) or 1
rec = {"step": step, "reward": round(sum(e.reward for e in flat) / len(flat), 4),
"acc": round(sum(e.correct for e in flat) / len(flat), 4),
"grounded": round(sum(e.grounded and e.correct for e in flat) / len(flat), 4),
"invented": round(sum(e.invented for e in flat) / len(flat), 4),
# reward-hacking monitors for the invented penalty: abstaining on answerable tasks,
# and wrong answers copied from some tool output
"nf_answerable": round(sum(e.task.answer != "NOT_FOUND" and normalize(e.ws.submitted or "") == "not_found"
for e in flat) / len(flat), 4),
"wrong_copied": round(sum(not e.correct and not e.invented and e.ws.submitted is not None
and normalize(e.ws.submitted) not in ("done", "not_found") for e in flat) / len(flat), 4),
"gen_tokens": round(sum(e.gen_tokens for e in flat) / len(flat), 1),
"tool_tokens": round(sum(e.tool_tokens for e in flat) / len(flat), 1),
"turns": round(sum(e.turns for e in flat) / len(flat), 2),
"flagged_turns": round(sum(sum(e.flags) for e in flat) / n_turns, 4),
"repeat_calls": round(sum(e.repeats for e in flat) / len(flat), 3),
"parse_err": round(sum(e.parse_errors for e in flat) / len(flat), 3),
"truncated": round(sum(e.truncated for e in flat) / len(flat), 3),
"groups": cnt["groups"] - c0.get("groups", 0),
"no_signal": cnt["no_signal"] - c0.get("no_signal", 0),
"stale_dropped": cnt["stale"] - c0.get("stale", 0),
"staleness": round(sum(roller.version - 1 - e.version for e in flat) / len(flat), 2),
"kind_acc": {k: round(sum(v) / len(v), 2) for k, v in sorted(per_kind.items())},
"succ_ema": {k: round(v, 2) for k, v in sorted(succ.items())},
**st, "grad_norm": round(gnorm, 4),
"rollout_s": round(t_roll, 1), "step_s": round(time.time() - t0, 1)}
if hasattr(roller, "stats"):
rec["engine"] = dict(roller.stats)
write(rec)
if step % a.eval_every == 0 or step == a.steps:
do_eval(step)
sd = {k: v.to(torch.bfloat16) if v.is_floating_point() else v for k, v in model.state_dict().items()}
torch.save({"model": sd, "config": asdict(model.cfg)}, os.path.join(a.out, f"rl_step{step}.pt"))
if __name__ == "__main__":
main()