Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
test_killswitch_integration.py371 linesDownload Raw Back to tests
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