tiny-agent-112m / code /tests /test_rollout.py
darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
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
@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