File size: 5,745 Bytes
c67f1b7 | 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 | """Tests for the training script's staging, curricula and checkpoint logic.
Separate from ``test_planner.py``, which covers the recursion itself. These
cover the wiring around it — the places where a run can be configured into
something other than what the stage table says, silently.
"""
import sys
from pathlib import Path
import pytest
import torch
sys.path.insert(0, str(Path(__file__).resolve().parents[2]))
from lejepa_control.losses import arrival_hold_loss # noqa: E402
from lejepa_control_2.scripts.train_planner import ( # noqa: E402
STAGES,
current_value,
parse_args,
parse_curriculum,
)
def test_stage_presets_match_the_bring_up_table():
"""Each stage turns on exactly one more mechanism than the last."""
a = parse_args(['--stage', 'A'])
assert a.cycles == 1 and a.use_feedback is False
assert (a.lambda_cycle, a.lambda_anchor, a.lambda_sat, a.lambda_support) \
== (0.0, 0.0, 0.0, 0.0)
b = parse_args(['--stage', 'B'])
assert b.cycles == 3 and b.use_feedback is True and b.lambda_cycle == 0.3
assert (b.lambda_anchor, b.lambda_sat, b.lambda_support) == (0.0, 0.0, 0.0)
c = parse_args(['--stage', 'C'])
assert (c.lambda_anchor, c.lambda_sat) == (0.05, 1e-3)
assert c.lambda_support == 0.0
d = parse_args(['--stage', 'D'])
assert d.lambda_support == 0.01
assert max(v for _, v in parse_curriculum(d.horizon_curriculum, 100)) == 5
# E is D with warm start off — the cold-start control, not a new mechanism
e = parse_args(['--stage', 'E'])
assert e.warm_start is False
for key in ('cycles', 'lambda_cycle', 'lambda_support', 'lambda_anchor',
'lambda_sat', 'horizon_curriculum'):
assert getattr(e, key) == getattr(d, key), f'E differs from D in {key}'
def test_explicit_flags_override_the_preset():
"""A preset must not silently win over something typed on the CLI."""
args = parse_args(['--stage', 'A', '--cycles', '5', '--lambda-cycle', '0.7'])
assert args.cycles == 5 and args.lambda_cycle == 0.7
# and untouched keys still come from the preset
assert args.use_feedback is False
def test_default_run_uses_the_section_10_hyperparameters():
args = parse_args([])
assert (args.inner, args.cycles) == (6, 3)
assert (args.hold_weight, args.alpha) == (0.5, 0.05)
assert (args.lambda_cycle, args.lambda_support) == (0.3, 0.01)
assert (args.lambda_anchor, args.lambda_sat) == (0.05, 1e-3)
assert (args.lr, args.weight_decay, args.grad_clip) == (1e-4, 1e-4, 1.0)
assert (args.batch_size, args.steps) == (32, 20000)
assert args.curriculum == '0:2,0.25:3,0.5:5'
assert parse_curriculum(args.horizon_curriculum, 20000) == [
(0, 3), (10000, 5)
]
def test_curriculum_advances_at_the_right_steps():
stages = parse_curriculum('0:2,0.25:3,0.5:5', 20000)
assert stages == [(0, 2), (5000, 3), (10000, 5)]
assert current_value(stages, 0) == 2
assert current_value(stages, 4999) == 2
assert current_value(stages, 5000) == 3
assert current_value(stages, 19999) == 5
@pytest.mark.parametrize(
'horizon_spec,step,want_h,want_q',
[
# stage D: the two curricula advance together, so the clamp is inert
('0:3,0.5:5', 0, 3, 2),
('0:3,0.5:5', 6000, 3, 3),
('0:3,0.5:5', 12000, 5, 5),
# stages B and C: H is pinned at 3 while the offset curriculum still
# climbs to 5. This is where the clamp actually bites.
('0:3', 12000, 3, 3),
],
)
def test_goal_offset_is_clamped_to_the_live_horizon(
horizon_spec, step, want_h, want_q
):
"""The offset curriculum must not outrun the horizon being rolled out.
Stages B and C hold ``H = 3`` while the shared offset curriculum advances
to 5 at the halfway point. Sampling ``q = 5`` against a 3-step rollout
would be silently clamped down to 3 inside ``arrival_hold_loss``, so the
deadline would stop meaning what the curriculum says it means — and the
dataset would be relabeling goals from 5 transitions ahead that the
planner has no steps left to reach. The clamp belongs at the sampler.
"""
horizon = current_value(parse_curriculum(horizon_spec, 20000), step)
offset_stages = parse_curriculum('0:2,0.25:3,0.5:5', 20000)
offset = min(current_value(offset_stages, step), horizon)
assert (horizon, offset) == (want_h, want_q)
assert offset <= horizon
def test_arrival_term_is_undistorted_once_the_offset_is_clamped():
"""With the clamp in place ``q`` always indexes a step that exists."""
d = torch.tensor([[0.9, 0.6, 0.3]])
for q in (1, 2, 3):
offset = torch.tensor([q])
got = arrival_hold_loss(d, offset, 0.5).mean()
hold = d[:, q:].mean(dim=1) if q < 3 else torch.zeros(1)
assert torch.allclose(got, (d[:, q - 1] + 0.5 * hold).mean())
def test_ablation_switches_are_reachable_but_off_by_default():
"""Section 12's ablations are wired behind flags, not on by default."""
args = parse_args([])
assert args.detach_schedule == 'last-cycle' # ablation 4
assert args.terminal_only is False # ablation 5
assert args.path_weighting == 'late' # ablation 6
assert args.use_feedback is True # ablation 2
assert parse_args(['--terminal-only']).terminal_only is True
assert parse_args(['--no-feedback']).use_feedback is False
assert parse_args(['--path-weighting', 'discount']).path_weighting \
== 'discount'
for schedule in ('last-cycle', 'one-step', 'full'):
assert parse_args(['--detach-schedule', schedule]).detach_schedule \
== schedule
if __name__ == '__main__':
sys.exit(pytest.main([__file__, '-v', '--no-header']))
|