File size: 8,508 Bytes
4397e12 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 | """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
|