Codeseys/composer-replication-framework
0
1"""R1 — HeldOutGuard wired into ComposerReplicationTrainer (the #2 safeguard).2 3These tests close the "tripwire exists but never fires" gap: the run-level4collapse kill-switch (``composer_replication.safety.HeldOutGuard``) must actually5be folded by the trainer at its logging cadence, fed the in-loop GRPO reward and6an INJECTED held-out eval, and halt the run on a fired verdict.7 8Acceptance gates:9 1. BACKWARD-COMPAT: no ``heldout_guard`` => ``_maybe_update_killswitch`` is a10 pure no-op (never touches the eval fn, never logs) — identical behavior.11 2. THE WIRING: a fake ``heldout_eval_fn`` that DECLINES while the in-loop12 reward RISES drives the guard to fire; strict mode raises13 ``CollapseStopError`` and the verdict carries the reward-hacking signature.14 3. Soft stop: ``strict_killswitch=False`` => no raise, the HF loop is asked to15 stop (``control.should_training_stop``).16 4. Healthy run (held-out tracks reward) never fires.17 5. KL-to-init is read from TRL's logged metric and reaches the guard.18 6. Constructor contract: ``heldout_guard`` without ``heldout_eval_fn`` raises.19 7. ``_compute_loss`` only folds the guard at the ``logging_steps`` cadence20 (not every micro-step), mirroring the loss-component logging.21 22CPU-only, no model download, no full GRPOTrainer init (stub instance via23``__new__`` + manual attribute wiring, the pattern used by the SDPO tests).24"""25from __future__ import annotations26 27import pytest28 29from composer_replication.safety import CollapseStopError, HeldOutGuard30from composer_replication.trainer.composer_trainer import ComposerReplicationTrainer31 32# ---------------------------------------------------------------------------33# Stubs — mirror _make_sdpo_trainer in test_sdpo_alignment_indices.py: build the34# trainer via __new__ so we never run GRPOTrainer.__init__ (no TRL setup, no35# model download), then wire only the attributes the kill-switch path reads.36# ---------------------------------------------------------------------------37 38class _State:39 def __init__(self, global_step: int = 0) -> None:40 self.global_step = global_step41 42 43class _Args:44 def __init__(self, logging_steps: int = 1) -> None:45 self.logging_steps = logging_steps46 47 48class _Control:49 def __init__(self) -> None:50 self.should_training_stop = False51 52 53def _make_killswitch_trainer(54 guard: HeldOutGuard | None,55 eval_fn,56 *,57 strict: bool = True,58 reward: float | None = 0.40,59 kl: float | None = None,60):61 """A ComposerReplicationTrainer stub exposing only what the kill-switch reads.62 63 ``reward`` / ``kl`` seed TRL's per-step metric series (the trainer reads the64 tail of ``self._metrics["train"][name]``). Pass reward=None to simulate "no65 reward aggregated yet".66 """67 obj = ComposerReplicationTrainer.__new__(ComposerReplicationTrainer)68 obj.heldout_guard = guard69 obj.heldout_eval_fn = eval_fn70 obj.strict_killswitch = strict71 obj.state = _State(global_step=0)72 obj.args = _Args(logging_steps=1)73 obj.control = _Control()74 train_metrics: dict[str, list] = {}75 if reward is not None:76 train_metrics["reward"] = [reward]77 if kl is not None:78 train_metrics["kl"] = [kl]79 obj._metrics = {"train": train_metrics}80 obj.logged: list[dict] = []81 # capture self.log(...) instead of routing through HF Trainer.log82 obj.log = obj.logged.append # type: ignore[assignment]83 return obj84 85 86def _set_step_reward(obj, step: int, reward: float, kl: float | None = None) -> None:87 obj.state.global_step = step88 obj._metrics["train"].setdefault("reward", []).append(reward)89 if kl is not None:90 obj._metrics["train"].setdefault("kl", []).append(kl)91 92 93# ---------------------------------------------------------------------------94# Gate 1 — backward compatibility: absent guard is a pure no-op95# ---------------------------------------------------------------------------96 97def test_absent_guard_is_noop():98 """No heldout_guard => the kill-switch path does nothing: the eval fn is99 never called, nothing is logged, no exception. This is the backward-compat100 guarantee (no kwarg => identical behavior)."""101 calls = {"n": 0}102 103 def eval_fn() -> float:104 calls["n"] += 1105 return 0.0106 107 # Pass eval_fn but NO guard — the helper must never reach the eval fn.108 obj = _make_killswitch_trainer(guard=None, eval_fn=eval_fn)109 for step in range(50):110 _set_step_reward(obj, step, reward=0.40 + 0.05 * step)111 obj._maybe_update_killswitch() # must be a no-op112 113 assert calls["n"] == 0, "held-out eval fn was called even though no guard set"114 assert obj.logged == [], "kill-switch logged even though no guard configured"115 assert obj.control.should_training_stop is False116 117 118def test_constructor_defaults_leave_killswitch_off():119 """The constructor defaults: a trainer built without the kwargs has120 heldout_guard=None / strict_killswitch defaulting on but inert."""121 obj = ComposerReplicationTrainer.__new__(ComposerReplicationTrainer)122 # Simulate the default-kwarg assignment __init__ performs.123 obj.heldout_guard = None124 obj.heldout_eval_fn = None125 obj.strict_killswitch = True126 obj._maybe_update_killswitch() # no state needed: returns immediately on None127 128 129# ---------------------------------------------------------------------------130# Gate 2 — THE WIRING: declining held-out + rising reward => guard fires & raises131# ---------------------------------------------------------------------------132 133def test_guard_fires_and_raises_on_reward_hacking_signature():134 """Fake heldout_eval_fn DECLINES while the in-loop reward RISES — the135 canonical reward-hacking signature. The wired guard must fire and (strict136 mode) raise CollapseStopError with the reward-hacking reason."""137 # min_steps small so the test is fast; isolate the decline-streak path.138 guard = HeldOutGuard(139 min_steps=3, decline_patience=3, ema_alpha=0.5, max_proxy_real_gap=10.0140 )141 142 # held-out declines every call; reward (TRL metric) rises every step.143 heldout = {"v": 0.80}144 145 def declining_eval() -> float:146 heldout["v"] -= 0.05147 return heldout["v"]148 149 obj = _make_killswitch_trainer(guard, declining_eval, strict=True)150 151 raised = None152 for step in range(1, 30):153 # rising in-loop reward fed via the TRL metric tail154 _set_step_reward(obj, step, reward=0.30 + 0.03 * step)155 try:156 obj._maybe_update_killswitch()157 except CollapseStopError as exc:158 raised = exc159 break160 161 assert raised is not None, "guard never fired on the reward-hacking signature"162 assert guard.should_halt()163 assert raised.status.fire164 # The fired verdict must be the held-out-declines-while-reward-rises signature.165 assert "held-out" in raised.status.reason166 assert raised.status.proxy_real_gap > 0.0 # proxy gained while real lost167 # And the kill-switch logged the verdict before raising.168 assert any("killswitch/fire" in d for d in obj.logged)169 170 171def test_soft_stop_sets_control_instead_of_raising():172 """strict_killswitch=False => a fired verdict does NOT raise; it sets the HF173 loop's control.should_training_stop so training ends gracefully."""174 guard = HeldOutGuard(175 min_steps=3, decline_patience=3, ema_alpha=0.5, max_proxy_real_gap=10.0176 )177 heldout = {"v": 0.80}178 179 def declining_eval() -> float:180 heldout["v"] -= 0.05181 return heldout["v"]182 183 obj = _make_killswitch_trainer(guard, declining_eval, strict=False)184 185 for step in range(1, 30):186 _set_step_reward(obj, step, reward=0.30 + 0.03 * step)187 obj._maybe_update_killswitch() # must NOT raise188 if obj.control.should_training_stop:189 break190 191 assert obj.control.should_training_stop is True, (192 "soft-stop guard fired but did not request training stop"193 )194 assert guard.should_halt()195 196 197# ---------------------------------------------------------------------------198# Gate 4 — healthy run never fires199# ---------------------------------------------------------------------------200 201def test_healthy_run_never_fires():202 """Held-out tracks the in-loop reward (both rise together), KL in band =>203 the wired guard never fires and training is never asked to stop."""204 guard = HeldOutGuard(205 min_steps=3, decline_patience=3, ema_alpha=0.5, kl_hard_stop=0.08,206 max_proxy_real_gap=10.0,207 )208 heldout = {"v": 0.28}209 210 def rising_eval() -> float:211 heldout["v"] += 0.01212 return heldout["v"]213 214 obj = _make_killswitch_trainer(guard, rising_eval, strict=True, kl=0.03)215 216 for step in range(1, 40):217 _set_step_reward(obj, step, reward=0.30 + 0.01 * step, kl=0.03)218 obj._maybe_update_killswitch() # must never raise on a healthy run219 220 assert not guard.should_halt()221 assert obj.control.should_training_stop is False222 223 224# ---------------------------------------------------------------------------225# Gate 5 — KL-to-init from TRL's logged metric reaches the guard226# ---------------------------------------------------------------------------227 228def test_kl_to_init_is_forwarded_to_guard():229 """The KL the trainer reads from TRL's "kl" metric must reach the guard's230 kl_ema (proves kl_to_init wiring), and a KL breach fires via the KL path."""231 guard = HeldOutGuard(232 min_steps=3, decline_patience=100, ema_alpha=0.5, kl_hard_stop=0.08,233 max_proxy_real_gap=10.0, # isolate the KL path234 )235 236 obj = _make_killswitch_trainer(guard, lambda: 0.40, strict=True, kl=0.04)237 238 # Warm-up with healthy KL; metrics flat so only the KL path can fire.239 for step in range(1, 5):240 _set_step_reward(obj, step, reward=0.40, kl=0.04)241 obj._maybe_update_killswitch()242 assert guard.last_status is not None and guard.last_status.kl_ema is not None, (243 "kl_to_init never reached the guard — KL wiring is broken"244 )245 246 # KL spikes above the hard stop; EMA climbs and crosses => fire via KL path.247 raised = None248 for step in range(5, 20):249 _set_step_reward(obj, step, reward=0.40, kl=0.30)250 try:251 obj._maybe_update_killswitch()252 except CollapseStopError as exc:253 raised = exc254 break255 assert raised is not None, "KL hard-stop never fired through the wired guard"256 assert "kl_to_init" in raised.status.reason257 258 259def test_no_reward_metric_yet_skips_cleanly():260 """Before TRL has aggregated any reward (empty metric series), the helper261 skips the fold rather than feeding a fabricated 0.0 into the guard's EMA."""262 guard = HeldOutGuard(min_steps=3, ema_alpha=0.5)263 calls = {"n": 0}264 265 def eval_fn() -> float:266 calls["n"] += 1267 return 0.40268 269 obj = _make_killswitch_trainer(guard, eval_fn, reward=None)270 obj._maybe_update_killswitch() # no reward series => skip271 272 assert calls["n"] == 0, "eval fn called despite no in-loop reward yet"273 assert guard.last_status is None, "guard advanced despite no reward metric"274 275 276# ---------------------------------------------------------------------------277# Gate 6 — constructor contract278# ---------------------------------------------------------------------------279 280def test_guard_without_eval_fn_raises_at_construction(monkeypatch):281 """A guard with no held-out eval is meaningless (the tripwire needs the282 held-out signal) => the REAL __init__ must reject it loudly. We stub the283 GRPOTrainer parent __init__ so the validation clause runs without a full284 TRL/model setup."""285 parent = ComposerReplicationTrainer.__bases__[0]286 monkeypatch.setattr(parent, "__init__", lambda self, *a, **k: None, raising=False)287 288 guard = HeldOutGuard(min_steps=3)289 with pytest.raises(ValueError, match="heldout_eval_fn"):290 ComposerReplicationTrainer(heldout_guard=guard) # no heldout_eval_fn291 292 293def test_guard_with_eval_fn_constructs_and_stays_off_when_absent(monkeypatch):294 """The real __init__ wires the kill-switch attributes; with both provided it295 constructs cleanly, and with neither provided the guard stays None (the296 default = OFF backward-compat path)."""297 parent = ComposerReplicationTrainer.__bases__[0]298 monkeypatch.setattr(parent, "__init__", lambda self, *a, **k: None, raising=False)299 300 # Both provided => constructs, guard wired.301 guard = HeldOutGuard(min_steps=3)302 t = ComposerReplicationTrainer(heldout_guard=guard, heldout_eval_fn=lambda: 0.4)303 assert t.heldout_guard is guard304 assert t.strict_killswitch is True # strict default305 306 # Neither provided => guard stays None (OFF).307 t2 = ComposerReplicationTrainer()308 assert t2.heldout_guard is None309 assert t2.heldout_eval_fn is None310 311 312# ---------------------------------------------------------------------------313# Gate 7 — _compute_loss only folds the guard at the logging cadence314# ---------------------------------------------------------------------------315 316def test_compute_loss_folds_guard_only_at_logging_cadence(monkeypatch):317 """Drive _compute_loss end-to-end (with the GRPO parent loss + SDPO/replay318 channels stubbed) and assert the guard is folded ONLY on logging-cadence319 steps — i.e. the kill-switch fold sits inside the same cadence gate as the320 loss-component logging, not on every micro-step."""321 import torch322 323 folds = {"n": 0}324 guard = HeldOutGuard(min_steps=10_000, ema_alpha=0.5) # never fires in this test325 326 def counting_eval() -> float:327 folds["n"] += 1328 return 0.40329 330 obj = ComposerReplicationTrainer.__new__(ComposerReplicationTrainer)331 obj.alpha_sdpo = 0.0332 obj.beta_replay = 0.0333 obj.heldout_guard = guard334 obj.heldout_eval_fn = counting_eval335 obj.strict_killswitch = True336 obj.state = _State(global_step=0)337 obj.args = _Args(logging_steps=10)338 obj.control = _Control()339 obj._metrics = {"train": {"reward": [0.40]}}340 obj.logged = []341 obj.log = obj.logged.append # type: ignore[assignment]342 343 # Stub the GRPO parent loss (the real `super()._compute_loss` would need a344 # full TRL trainer) and the SDPO / replay channels to zero. We patch the345 # PARENT class's _compute_loss so `super()._compute_loss(...)` resolves to it.346 parent = ComposerReplicationTrainer.__bases__[0]347 monkeypatch.setattr(348 parent, "_compute_loss",349 lambda self, model, inputs: torch.tensor(1.0),350 raising=False,351 )352 monkeypatch.setattr(353 ComposerReplicationTrainer, "_compute_sdpo_loss",354 lambda self, model, inputs: torch.tensor(0.0),355 raising=True,356 )357 monkeypatch.setattr(358 ComposerReplicationTrainer, "_compute_trace_replay_loss",359 lambda self, model, inputs: torch.tensor(0.0),360 raising=True,361 )362 363 for step in range(0, 35):364 obj.state.global_step = step365 obj._metrics["train"]["reward"].append(0.40 + 0.001 * step)366 total = obj._compute_loss(model=object(), inputs={})367 assert float(total.detach()) == pytest.approx(1.0)368 369 # steps 0, 10, 20, 30 are the only cadence hits => exactly 4 guard folds.370 assert folds["n"] == 4, f"expected 4 cadence folds, got {folds['n']}"371 