File size: 4,007 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
"""Reward shaping and credit assignment math (tiny_agent.rl)."""
import random

import pytest

from tiny_agent.rl import (LengthPenalty, base_reward, group_advantages, has_signal, length_deltas,
                           signed_rebalance)


def sig(turns, inp, out):
    return {"turns": turns, "input_tokens": inp, "output_tokens": out}


def test_base_reward():
    assert base_reward(True, True) == 1.0
    assert base_reward(True, False) == 0.5
    assert base_reward(False, True) == 0.0


def test_length_penalty_only_on_passes_and_gated():
    # 5 of 8 pass (> 50%): the long pass is penalized, failures never are
    r = [1, 1, 1, 1, 1, 0, 0, 0]
    s = [sig(3, 100, 100)] * 4 + [sig(3, 100, 400)] + [sig(8, 900, 900)] * 3
    d = length_deltas(r, s)
    assert d[:4] == [0.0] * 4 and d[5:] == [0.0] * 3
    assert d[4] == pytest.approx(-0.2)          # 4x the anchor saturates
    # exactly half passing: no penalty at all (strictly greater required)
    assert length_deltas([1, 1, 0, 0], [sig(1, 1, 1), sig(5, 5, 50), sig(1, 1, 1), sig(1, 1, 1)]) == [0.0] * 4


def test_length_penalty_ramp():
    # anchor = p30 of passes; 1.5x the anchor on the worst metric -> 0.2 * 0.5**1.5
    r = [1.0] * 4
    s = [sig(2, 100, 100), sig(2, 100, 100), sig(2, 100, 100), sig(3, 100, 100)]
    d = length_deltas(r, s)
    assert d[3] == pytest.approx(-0.2 * 0.5 ** 1.5)
    assert d[:3] == [0.0] * 3


def test_dynamic_sampling_signal():
    assert not has_signal([1, 1, 1])
    assert not has_signal([0, 0, 0])
    assert has_signal([1, 0.8, 1])
    assert group_advantages([1, 0, 0, 1]) == [0.5, -0.5, -0.5, 0.5]


def test_signed_rebalance_conserves_mass():
    rng = random.Random(0)
    for _ in range(200):
        n = 16
        adv = [rng.choice([-1, 1]) * rng.random() for _ in range(n)]
        nf = [rng.randrange(0, 20) for _ in range(n)]
        nc = [rng.randrange(1, 200) for _ in range(n)]
        w = signed_rebalance(adv, nf, nc, max_scale=1e9, min_scale=-1e9)   # unclipped
        pos_before = sum(a * (f + c) for a, f, c in zip(adv, nf, nc) if a > 0)
        pos_after = sum(a * (wc * c + wf * f) for a, f, c, (wc, wf) in zip(adv, nf, nc, w) if a > 0)
        neg_before = sum(a * (f + c) for a, f, c in zip(adv, nf, nc) if a < 0)
        neg_after = sum(a * (wc * c + wf * f) for a, f, c, (wc, wf) in zip(adv, nf, nc, w) if a < 0)
        assert pos_after == pytest.approx(pos_before)
        assert neg_after == pytest.approx(neg_before)


def test_signed_rebalance_directions():
    w = signed_rebalance([0.5, -0.5, 0.0], [10, 10, 10], [100, 100, 100])
    (pc, pf), (nc, nf), (zc, zf) = w
    assert pf == 0.0 and pc > 1.0      # flagged tokens of a good episode get no credit
    assert nf == 2.0 and nc < 1.0      # flagged tokens of a bad episode get double blame
    assert (zc, zf) == (0.0, 0.0)
    # clipping bounds hold
    w = signed_rebalance([1.0, -1.0], [1000, 1000], [1, 1])
    assert w[0][0] == 2.0 and w[1][0] == 0.5


def test_base_reward_invented():
    from tiny_agent.rl import base_reward
    assert base_reward(True, True) == 1.0 and base_reward(True, False) == 0.5
    assert base_reward(False, False) == 0.0
    assert base_reward(False, False, invented=True, invented_penalty=0.3) == -0.3
    assert base_reward(False, False, invented=True) == 0.0      # off by default


def test_invented_detection():
    import random
    from tiny_agent.tasks import make_task, invented, grounded
    t = make_task(random.Random(1), "config_value")
    msgs = [{"role": "tool", "results": [f"port: 4242\nowner: Ann Example\n"]}]
    assert invented(t, "Brodai", msgs)                 # appears nowhere
    assert not invented(t, "4242", msgs)               # read from a result
    assert not invented(t, "ann example", msgs)        # case-insensitive
    assert not invented(t, "NOT_FOUND", msgs) and not invented(t, "DONE", msgs) and not invented(t, None, msgs)
    assert invented(t, "424", msgs)                    # token boundaries: a prefix of a number is not seen