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