Download code/tests/test_rollout.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 8.51 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tests/test_rollout.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/tests/test_rollout.py
-
curl -L -o test_rollout.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tests/test_rollout.py
8.51 kB
| """Rollout schedulers driven by a scripted policy (the oracle's turns, token by token) on a small | |
| random model: episodes must run real tool calls to completion, and the log-probs recorded while | |
| decoding (with slot reuse, sub-batch prefill and tool feedback) must match a full forward pass.""" | |
| import random | |
| import pytest | |
| import torch | |
| import torch.nn.functional as F | |
| from tokenizers import Tokenizer | |
| from tests.test_model import DEV, small_cfg | |
| from tiny_agent.chat import render | |
| from tiny_agent.model import TinyAgentLM | |
| from tiny_agent.rollout import Engine, Roller | |
| from tiny_agent.tasks import TRAIN_KINDS, make_task, oracle_plan | |
| from tiny_agent.text import DATA | |
| TOK_PATH = f"{DATA}/tokenizer.json" | |
| def scripted(cls): | |
| class Scripted(cls): | |
| def _sample(self, logits, row_eps): | |
| toks = [] | |
| for e in row_eps: | |
| script = getattr(e, "_script", None) if e is not None else None | |
| toks.append(script[e.gen_tokens] if script and e.gen_tokens < len(script) else 0) | |
| tok = torch.tensor(toks, device=logits.device) | |
| lp = F.log_softmax(logits.float(), -1).gather(1, tok[:, None]).squeeze(1) | |
| return tok, toks, lp.tolist() | |
| return Scripted | |
| def make_script(tok, task, rng): | |
| ids = [] | |
| prefix = len("<|im_start|>assistant\n") | |
| for thought, calls in oracle_plan(task, rng): | |
| text = render([{"role": "assistant", "think": thought, "content": "", "tool_calls": calls}]) | |
| ids += tok.encode(text[prefix:].rstrip("\n"), add_special_tokens=False).ids | |
| return ids | |
| def setup(): | |
| tok = Tokenizer.from_file(TOK_PATH) | |
| torch.manual_seed(0) | |
| m = TinyAgentLM(small_cfg(vocab_size=tok.get_vocab_size(), max_seq_len=4096, d_model=64, | |
| n_heads=2, head_dim=32, n_kv_heads=1)).to(DEV).float().eval() | |
| with torch.no_grad(): | |
| for p in m.parameters(): | |
| p.normal_(0, 0.05) | |
| return tok, m | |
| def run_scripted(cls, tok, m, n, **kw): | |
| tasks = [make_task(random.Random(7_000 + i), TRAIN_KINDS[i % len(TRAIN_KINDS)]) for i in range(n)] | |
| eng = scripted(cls)(m, tok, device=DEV, max_len=4096, **kw) | |
| if cls is Engine: | |
| eng.cache = type(eng.cache)(m, eng.B, 4096, DEV, dtype=torch.float32) | |
| scripts = [make_script(tok, t, random.Random(7_000 + i)) for i, t in enumerate(tasks)] | |
| if cls is Roller: | |
| # Roller creates its episodes inside run(): attach the scripts on the first prefill | |
| eps_holder = {} | |
| orig_feed = eng._feed | |
| def feed(cache, chunks, eps): | |
| if "eps" not in eps_holder: | |
| for e, sc in zip(eps, scripts): | |
| e._script = sc | |
| eps_holder["eps"] = eps | |
| return orig_feed(cache, chunks, eps) | |
| eng._feed = feed | |
| return eng.run(tasks), tasks | |
| eps = eng.submit(tasks) | |
| for e, sc in zip(eps, scripts): | |
| e._script = sc | |
| while any(not e.done for e in eps): | |
| eng.tick() | |
| return eps, tasks | |
| def check_logp(m, eps): | |
| for e in eps: | |
| x = torch.tensor(e.tokens, device=DEV)[None] | |
| with torch.no_grad(): | |
| logits = m(x, torch.zeros_like(x))[0].float() | |
| lp = F.log_softmax(logits[:-1], -1).gather(1, x[0, 1:, None]).squeeze(1) | |
| gm = torch.tensor(e.gen_mask[1:], device=DEV).bool() | |
| rec = torch.tensor(e.logp[1:], device=DEV) | |
| err = (lp[gm] - rec[gm]).abs().max().item() | |
| assert err < 2e-3, err | |
| def test_engine_scripted_episodes(setup): | |
| tok, m = setup | |
| # 3 slots for 9 episodes: slots are reused; prefill batching thresholds are exercised | |
| eps, tasks = run_scripted(Engine, tok, m, 9, slots=3, prefill_rows=2, prefill_wait=3, max_turn_tokens=400) | |
| assert [e.idx for e in eps] == list(range(9)) | |
| assert sum(e.correct for e in eps) >= 8, [(t.kind, e.correct, e.ws.submitted) for t, e in zip(tasks, eps)] | |
| for e in eps: | |
| assert len(e.turn_spans) == e.turns == len(e.flags) == sum(m_["role"] == "assistant" for m_ in e.messages) | |
| assert not any(e.flags) and e.repeats == 0 | |
| assert sum(e.gen_mask) == e.gen_tokens and len(e.logp) == len(e.tokens) == len(e.gen_mask) | |
| for a, b in e.turn_spans: | |
| assert all(e.gen_mask[a:b]) | |
| check_logp(m, eps) | |
| def test_lockstep_matches_engine(setup): | |
| tok, m = setup | |
| a, _ = run_scripted(Engine, tok, m, 4, slots=2, prefill_rows=1, prefill_wait=1, max_turn_tokens=400) | |
| b, _ = run_scripted(Roller, tok, m, 4, max_turn_tokens=400) | |
| for x, y in zip(a, b): | |
| assert x.tokens == y.tokens and x.gen_mask == y.gen_mask and x.correct == y.correct | |
| assert x.turn_spans == y.turn_spans | |
| def test_finalize_flags_bad_turns(): | |
| from tiny_agent.rollout import Episode, finalize | |
| from tiny_agent.tools import Workspace | |
| t = make_task(random.Random(3), "config_value") | |
| ls = {"name": "bash", "arguments": {"command": "ls"}} | |
| msgs = t.messages() + [ | |
| {"role": "assistant", "think": None, "content": "", "tool_calls": [ls]}, | |
| {"role": "tool", "results": ["a"]}, | |
| {"role": "assistant", "think": None, "content": "", "tool_calls": [{"error": "Error: could not parse"}]}, | |
| {"role": "tool", "results": ["Error: could not parse"]}, | |
| {"role": "assistant", "think": None, "content": "", "tool_calls": [{"name": "grep", "arguments": {}}]}, | |
| {"role": "tool", "results": ["Error: unknown tool 'grep'. Available: ..."]}, | |
| {"role": "assistant", "think": None, "content": "", "tool_calls": [ls]}, | |
| {"role": "tool", "results": ["a"]}, | |
| ] | |
| e = finalize(Episode(t, Workspace(t.files), msgs)) | |
| assert e.flags == [False, True, True, True] and e.repeats == 1 and not e.correct | |
| def test_policy_gradient_direction(setup): | |
| """One SGD step on the GRPO loss raises log-probs of the positive episode's tokens and lowers | |
| the negative one's; on fresh rollouts the importance ratio is ~1 (no train/decode mismatch).""" | |
| from scripts.grpo import policy_loss, token_weights | |
| tok, m = setup | |
| eps, _ = run_scripted(Engine, tok, m, 2, slots=2, prefill_rows=1, prefill_wait=1, max_turn_tokens=400) | |
| eps[0].reward, eps[1].reward = 1.0, 0.0 | |
| items = token_weights([eps], kappa=2.0, min_scale=0.5, max_scale=2.0) | |
| def mean_lp(e): | |
| x = torch.tensor(e.tokens, device=DEV)[None] | |
| with torch.no_grad(): | |
| lg = m(x, torch.zeros_like(x))[0].float() | |
| lp = F.log_softmax(lg[:-1], -1).gather(1, x[0, 1:, None]).squeeze(1) | |
| g = torch.tensor(e.gen_mask[1:], device=DEV).bool() | |
| return lp[g].mean().item() | |
| before = [mean_lp(e) for e in eps] | |
| m.train() | |
| st = policy_loss(m, items, micro=2, device=DEV, T=4096, clip_lo=0.2, clip_hi=5.0) | |
| assert st["mismatch"] < 0.05 and st["clip_frac"] == 0.0, st | |
| with torch.no_grad(): | |
| for p in m.parameters(): | |
| if p.grad is not None: | |
| p -= 0.5 * p.grad | |
| p.grad = None | |
| m.eval() | |
| after = [mean_lp(e) for e in eps] | |
| assert after[0] > before[0] and after[1] < before[1], (before, after) | |
| def test_junk_arguments_never_raise(): | |
| """Sampled policies produce any JSON; tools must answer with an error result, never raise, and | |
| scoring must survive whatever the model left in the workspace.""" | |
| from tiny_agent.rollout import Episode, finalize | |
| from tiny_agent.tools import Workspace | |
| junk = [None, 5, -1, 1e300, "ten", [], {}, ["a"], True, "../../etc/passwd", "", "\x00"] | |
| t = make_task(random.Random(5), "edit") | |
| with Workspace(t.files) as ws: | |
| for name in ("bash", "read", "edit", "write", "submit", "nope"): | |
| for v in junk: | |
| for args in ({"path": v}, {"command": v}, {"command": "ls", "timeout": v}, | |
| {"path": "README.md", "offset": v}, {"path": "README.md", "limit": v}, | |
| {"path": v, "old_string": v, "new_string": v}, {"path": v, "content": v}, | |
| {"answer": v}, {}, {"zzz": 1}): | |
| assert isinstance(ws.call(name, args), str) | |
| for kind in ("edit", "write_fact", "fix_test"): | |
| t = make_task(random.Random(6), kind) | |
| ws = Workspace(t.files) | |
| for rel in t.files: # corrupt every file the checker might parse | |
| ws.call("write", {"path": rel, "content": "{{{: [unclosed\n\x00"}) | |
| ws.call("submit", {"answer": "DONE"}) | |
| e = finalize(Episode(t, ws, t.messages())) | |
| assert e.correct is False | |