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