Download code/scripts/grpo.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 14.8 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/grpo.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/grpo.py
-
curl -L -o grpo.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/grpo.py
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() | |