"""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 @pytest.fixture(scope="module") 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