"""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