Download code/tests/test_rl.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 4.01 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tests/test_rl.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/tests/test_rl.py
-
curl -L -o test_rl.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tests/test_rl.py
4.01 kB
| """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 | |