From 925e600702dc02277462ce85dea21aa9a0d36b12 Mon Sep 17 00:00:00 2001 From: Alpamys Date: Wed, 1 Jul 2026 15:05:32 +0500 Subject: [PATCH] feat(train): reversible bang-bang KL controller + hysteresis (v0.71.26 Part B) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Stage 1 of closed-loop reward-hacking mitigation: the first closed loop. - schema: 8 Stage-1 tunables (beta_floor/ceil, trip/release band, dwell_steps, release_patience, kl_gain, signals allowlist) + bounds; extended _validate_reward_hack_compat (floor 91). --- src/soup_cli/commands/train.py | 36 ++ src/soup_cli/config/schema.py | 147 +++++++++ src/soup_cli/trainer/ppo.py | 35 ++ src/soup_cli/utils/peft_wiring.py | 13 + src/soup_cli/utils/reward_hack_control.py | 250 +++++++++++++- tests/test_v07126.py | 379 ++++++++++++++++++++++ 6 files changed, 850 insertions(+), 10 deletions(-) diff --git a/src/soup_cli/commands/train.py b/src/soup_cli/commands/train.py index 62abb1a..5e1b0e9 100644 --- a/src/soup_cli/commands/train.py +++ b/src/soup_cli/commands/train.py @@ -167,6 +167,15 @@ def train( "training.echo_trap_enabled=true on grpo/ppo." ), ), + reward_hack_mitigation: str = typer.Option( + None, + "--reward-hack-mitigation", + help=( + "Closed-loop reward-hacking mitigation mode: off | log_only | " + "kl_control | pid_lagrangian. Requires training.reward_hack_detector " + "on grpo/ppo. Overrides training.reward_hack_mitigation. (v0.71.26)" + ), + ), minillm_on_policy: bool = typer.Option( False, "--minillm-on-policy", @@ -344,6 +353,29 @@ def train( cfg.training.echo_trap_tokenizer_aware = True console.print("[green]Echo-trap tokenizer-aware scoring enabled[/]") + # --- Reward-hack mitigation shortcut (v0.71.26) --- + if reward_hack_mitigation is not None: + valid_modes = ("off", "log_only", "kl_control", "pid_lagrangian") + if reward_hack_mitigation not in valid_modes: + console.print( + "[red]--reward-hack-mitigation must be one of " + f"{', '.join(valid_modes)}[/]" + ) + raise typer.Exit(1) + if ( + reward_hack_mitigation != "off" + and cfg.training.reward_hack_detector is None + ): + console.print( + "[red]--reward-hack-mitigation requires " + "training.reward_hack_detector to be set (the signal source)[/]" + ) + raise typer.Exit(1) + cfg.training.reward_hack_mitigation = reward_hack_mitigation + console.print( + f"[green]Reward-hack mitigation:[/] {reward_hack_mitigation}" + ) + # --- MiniLLM on-policy rollout shortcut (v0.71.18 #257) --- if minillm_on_policy: if not cfg.training.minillm_enabled: @@ -618,6 +650,10 @@ def train( script_args.append("--tensorboard") if echo_trap_tokenizer_aware: script_args.append("--echo-trap-tokenizer-aware") + if reward_hack_mitigation is not None: + script_args.extend( + ["--reward-hack-mitigation", reward_hack_mitigation] + ) if yes: script_args.append("--yes") argv = build_accelerate_argv( diff --git a/src/soup_cli/config/schema.py b/src/soup_cli/config/schema.py index 551cde2..4045026 100644 --- a/src/soup_cli/config/schema.py +++ b/src/soup_cli/config/schema.py @@ -1672,6 +1672,75 @@ class TrainingConfig(BaseModel): "{'grpo', 'ppo'} on a non-mlx backend." ), ) + reward_hack_beta_floor: float = Field( + default=0.02, + gt=0.0, + description=( + "v0.71.26 — lower β/kl_coef bound for the mitigation controller. " + "Must be > 0 (β=0 gates off the ref-log-prob path at generation)." + ), + ) + reward_hack_beta_ceil: float = Field( + default=1.0, + gt=0.0, + le=1000.0, + description=( + "v0.71.26 — upper β/kl_coef bound for the mitigation controller. " + "Must be > reward_hack_beta_floor." + ), + ) + reward_hack_trip_band: float = Field( + default=0.30, + ge=0.0, + le=1.0, + description=( + "v0.71.26 — hacking drop_pct at/above which the controller wants " + "to RAISE β. Must be > reward_hack_release_band." + ), + ) + reward_hack_release_band: float = Field( + default=0.10, + ge=0.0, + le=1.0, + description=( + "v0.71.26 — hacking drop_pct at/below which the controller wants " + "to RELAX β. Must be < reward_hack_trip_band." + ), + ) + reward_hack_dwell_steps: int = Field( + default=2, + ge=1, + le=100_000, + description=( + "v0.71.26 — consecutive trip-band steps required before the " + "controller raises β (hysteresis)." + ), + ) + reward_hack_release_patience: int = Field( + default=3, + ge=1, + le=100_000, + description=( + "v0.71.26 — consecutive release-band steps required before the " + "controller relaxes β." + ), + ) + reward_hack_kl_gain: float = Field( + default=1.5, + gt=1.0, + le=100.0, + description=( + "v0.71.26 — multiplicative β step per trip (>1). β is multiplied " + "on trip / divided on release, clamped to [floor, ceil]." + ), + ) + reward_hack_signals: List[str] = Field( + default_factory=lambda: ["info_rm"], + description=( + "v0.71.26 — signals combined into the controller's multi-signal " + "vote. Allowlist: info_rm, rm_ensemble, length_trend, repetition." + ), + ) @field_validator("reward_hack_mitigation", mode="before") @classmethod @@ -3073,6 +3142,72 @@ class EvalConfig(BaseModel): ) +# --- v0.71.26 reward-hack mitigation validation helpers --- + +# Control tunables that are meaningless unless a mitigation mode is set. Setting +# any to a non-default value while reward_hack_mitigation='off' is a silent +# no-op footgun (mirrors the v0.70.0 minillm offenders-list policy). Extended +# per stage (Stage 2/3 tunables added with their fields). +_REWARD_HACK_TUNABLE_DEFAULTS: dict = { + "reward_hack_beta_floor": 0.02, + "reward_hack_beta_ceil": 1.0, + "reward_hack_trip_band": 0.30, + "reward_hack_release_band": 0.10, + "reward_hack_dwell_steps": 2, + "reward_hack_release_patience": 3, + "reward_hack_kl_gain": 1.5, + "reward_hack_signals": ["info_rm"], +} + + +def _customized_reward_hack_tunables(tcfg) -> list: + """Return the reward-hack control tunables set to a non-default value.""" + offenders = [] + for field_name, default in _REWARD_HACK_TUNABLE_DEFAULTS.items(): + if getattr(tcfg, field_name, default) != default: + offenders.append(field_name) + return offenders + + +def _validate_reward_hack_controller(tcfg) -> None: + """Validate the mitigation-controller config (only when a mode is active). + + Numeric consistency (β floor < ceil, release < trip band), the signal + allowlist, and the β-schedule mutual exclusion. + """ + floor = tcfg.reward_hack_beta_floor + ceil = tcfg.reward_hack_beta_ceil + if floor >= ceil: + raise ValueError( + f"reward_hack_beta_floor ({floor}) must be < " + f"reward_hack_beta_ceil ({ceil})" + ) + release = tcfg.reward_hack_release_band + trip = tcfg.reward_hack_trip_band + if release >= trip: + raise ValueError( + f"reward_hack_release_band ({release}) must be < " + f"reward_hack_trip_band ({trip})" + ) + from soup_cli.utils.reward_hack_control import SIGNAL_NAMES + + for name in tcfg.reward_hack_signals or []: + if name not in SIGNAL_NAMES: + raise ValueError( + f"reward_hack_signals contains unknown signal {name!r}; " + f"valid: {sorted(SIGNAL_NAMES)}" + ) + # A control mode drives the KL/ref dynamics; a competing β schedule + # (ref_model_ema_alpha regenerates the reference) fights it — reject. + if tcfg.reward_hack_mitigation in ("kl_control", "pid_lagrangian"): + if getattr(tcfg, "ref_model_ema_alpha", None) is not None: + raise ValueError( + "reward_hack_mitigation kl_control/pid_lagrangian is mutually " + "exclusive with ref_model_ema_alpha (both drive the KL/ref " + "dynamics); pick one" + ) + + class SoupConfig(BaseModel): """Root config for soup.yaml.""" @@ -4643,6 +4778,14 @@ class SoupConfig(BaseModel): detector = tcfg.reward_hack_detector halt = tcfg.reward_hack_halt mitigation = getattr(tcfg, "reward_hack_mitigation", "off") + # v0.71.26 — footgun: control tunables set while mitigation is off. + if mitigation == "off": + offenders = _customized_reward_hack_tunables(tcfg) + if offenders: + raise ValueError( + f"reward-hack tunables {offenders} require " + "reward_hack_mitigation to be set (not 'off')" + ) if detector is None and not halt and mitigation == "off": return self # halt without detector is a silent no-op footgun — reject. @@ -4657,6 +4800,10 @@ class SoupConfig(BaseModel): f"reward_hack_mitigation={mitigation!r} requires " "reward_hack_detector to be set (the signal source)" ) + # v0.71.26 — controller config (numeric bounds, signal allowlist, + # β-schedule mutual exclusion) only when a mode is active. + if mitigation != "off": + _validate_reward_hack_controller(tcfg) if self.task not in ("grpo", "ppo"): raise ValueError( "reward_hack_detector / reward_hack_halt / " diff --git a/src/soup_cli/trainer/ppo.py b/src/soup_cli/trainer/ppo.py index 711e7e0..2dceebd 100644 --- a/src/soup_cli/trainer/ppo.py +++ b/src/soup_cli/trainer/ppo.py @@ -191,6 +191,28 @@ class PPOTrainerWrapper: if self.reward_fn is not None: reward_funcs.append(self.reward_fn) + # v0.71.26 — reward-hack mitigation buffer parity with GRPO. When a + # detector / mitigation mode is set, capture the callable reward fns' + # rewards + completions into a shared RLSignalBuffer so the mitigation + # callback can observe the step (BETA — the on-GPU proof is GRPO-only). + # nn.Module reward models are skipped (their call shape differs). + self._rl_buffer = None + from soup_cli.utils.peft_wiring import rl_callbacks_need_buffer + + if rl_callbacks_need_buffer(tcfg) and reward_funcs: + from soup_cli.utils.rl_signal_buffer import ( + RLSignalBuffer, + wrap_reward_funcs, + ) + + self._rl_buffer = RLSignalBuffer() + reward_funcs = [ + wrap_reward_funcs(fn, self._rl_buffer) + if callable(fn) and not hasattr(fn, "forward") + else fn + for fn in reward_funcs + ] + # --- Trainer --- ppo_trainer_params = inspect.signature(ppo_trainer_cls.__init__).parameters @@ -255,6 +277,19 @@ class PPOTrainerWrapper: self._dataset_in_constructor = True self.trainer = ppo_trainer_cls(**trainer_kwargs) + # v0.71.26 — reward-hack mitigation / echo-trap / RL-checkpoint callbacks + # (PPO parity with GRPO; kl_coef mutation for the controller). + from soup_cli.utils.peft_wiring import attach_rl_callbacks + + attach_rl_callbacks( + self.trainer, + tcfg, + buffer=self._rl_buffer, + tokenizer=self.tokenizer, + output_dir=str(output_dir), + task="ppo", + ) + # v0.40.6 #67 — ReLoRA callback (magnitude-prune LoRA every N steps). from soup_cli.utils.peft_wiring import ( attach_curriculum_callback, diff --git a/src/soup_cli/utils/peft_wiring.py b/src/soup_cli/utils/peft_wiring.py index 0661679..0497f17 100644 --- a/src/soup_cli/utils/peft_wiring.py +++ b/src/soup_cli/utils/peft_wiring.py @@ -233,6 +233,7 @@ def _attach_reward_hack( mitigation = getattr(tcfg, "reward_hack_mitigation", "off") if mitigation != "off" and detector is not None: from soup_cli.utils.reward_hack_control import ( + BangBangPolicy, MitigationLogWriter, RewardHackMitigationCallback, ) @@ -244,6 +245,17 @@ def _attach_reward_hack( signals = tuple( getattr(tcfg, "reward_hack_signals", None) or ("info_rm",) ) + bang_bang = None + if mitigation == "kl_control": + bang_bang = BangBangPolicy( + beta_floor=tcfg.reward_hack_beta_floor, + beta_ceil=tcfg.reward_hack_beta_ceil, + trip_band=tcfg.reward_hack_trip_band, + release_band=tcfg.reward_hack_release_band, + dwell_steps=tcfg.reward_hack_dwell_steps, + release_patience=tcfg.reward_hack_release_patience, + kl_gain=tcfg.reward_hack_kl_gain, + ) callback = RewardHackMitigationCallback( mode=mitigation, detector=detector, @@ -252,6 +264,7 @@ def _attach_reward_hack( buffer=buffer, tokenizer=tokenizer, task=task, + bang_bang=bang_bang, ) trainer.add_callback(callback) callback.attach(trainer) diff --git a/src/soup_cli/utils/reward_hack_control.py b/src/soup_cli/utils/reward_hack_control.py index a62262b..691ee46 100644 --- a/src/soup_cli/utils/reward_hack_control.py +++ b/src/soup_cli/utils/reward_hack_control.py @@ -36,7 +36,7 @@ import statistics import threading import time from collections.abc import Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, replace from pathlib import Path from typing import Any @@ -217,6 +217,147 @@ def mean_repetition(completions: Sequence[Any]) -> float: return 0.0 +# --- bang-bang controller (Stage 1) --- + + +@dataclass(frozen=True) +class MitigationAction: + """One controller step's decision (frozen). + + - ``new_beta``: the β / kl_coef the trainer should apply next step. + - ``tripped``: whether the controller is in the raised band. + - ``verdict``: OK / WARN / HACK from the classifier over the vote. + - ``reason``: human-readable summary for the mitigation log. + """ + + new_beta: float + tripped: bool + verdict: str + reason: str + + +@dataclass(frozen=True) +class BangBangPolicy: + """Reversible bang-bang + hysteresis controller policy. + + Two bands with a dead-zone between them: the controller wants to RAISE β + when the vote is at/above ``trip_band`` and RELAX when at/below + ``release_band``. ``dwell_steps`` consecutive want-raise steps are required + before the first trip, and ``release_patience`` consecutive want-relax + steps before relaxing — so a signal that flaps across the bands does not + flap β. β moves geometrically by ``kl_gain`` and is clamped to + ``[beta_floor, beta_ceil]`` (never crossing 0). + """ + + beta_floor: float + beta_ceil: float + trip_band: float + release_band: float + dwell_steps: int + release_patience: int + kl_gain: float + + def __post_init__(self) -> None: + _check_finite_float(self.beta_floor, "beta_floor", nonneg=True) + _check_finite_float(self.beta_ceil, "beta_ceil", nonneg=True) + _check_finite_float(self.trip_band, "trip_band", nonneg=True) + _check_finite_float(self.release_band, "release_band", nonneg=True) + _check_finite_float(self.kl_gain, "kl_gain", nonneg=True) + _check_nonneg_int(self.dwell_steps, "dwell_steps") + _check_nonneg_int(self.release_patience, "release_patience") + if self.beta_floor <= 0.0: + raise ValueError(f"beta_floor must be > 0, got {self.beta_floor}") + if self.beta_floor >= self.beta_ceil: + raise ValueError( + f"beta_floor ({self.beta_floor}) must be < " + f"beta_ceil ({self.beta_ceil})" + ) + if not 0.0 <= self.release_band < self.trip_band <= 1.0: + raise ValueError( + "require 0 <= release_band < trip_band <= 1, got " + f"release_band={self.release_band}, trip_band={self.trip_band}" + ) + if self.dwell_steps < 1: + raise ValueError("dwell_steps must be >= 1") + if self.release_patience < 1: + raise ValueError("release_patience must be >= 1") + if self.kl_gain <= 1.0: + raise ValueError(f"kl_gain must be > 1, got {self.kl_gain}") + + +def _verdict_for(vote: float) -> str: + from soup_cli.utils.reward_hacking import classify_hack_signal + + return classify_hack_signal(min(1.0, max(0.0, float(vote)))) + + +def bang_bang_step( + policy: BangBangPolicy, state: ControllerState, *, vote: float +) -> tuple[ControllerState, MitigationAction]: + """Advance the bang-bang controller one step given the multi-signal ``vote``. + + ``vote`` is the combined hacking signal in ``[0, 1]`` (see + :func:`combine_signals`). Returns the next :class:`ControllerState` and the + :class:`MitigationAction` (the β the trainer should apply). + """ + fvote = float(vote) + beta = state.beta if state.beta > 0.0 else policy.beta_floor + tripped = state.tripped + dwell = state.dwell_count + release = state.release_count + reason = "hold" + + if fvote >= policy.trip_band: + release = 0 + if not tripped: + dwell += 1 + if dwell >= policy.dwell_steps: + beta = min(policy.beta_ceil, beta * policy.kl_gain) + tripped = True + dwell = 0 + reason = f"trip: raise beta to {beta:.4f} (vote={fvote:.3f})" + else: + new_beta = min(policy.beta_ceil, beta * policy.kl_gain) + if new_beta != beta: + reason = f"raise beta to {new_beta:.4f} (vote={fvote:.3f})" + beta = new_beta + elif fvote <= policy.release_band: + dwell = 0 + if tripped: + release += 1 + if release >= policy.release_patience: + new_beta = max(policy.beta_floor, beta / policy.kl_gain) + if new_beta != beta: + reason = f"relax beta to {new_beta:.4f} (vote={fvote:.3f})" + beta = new_beta + if beta <= policy.beta_floor: + tripped = False + release = 0 + else: + release = 0 + else: + # dead-band: hold and decay both counters (hysteresis). + dwell = 0 + release = 0 + + new_state = replace( + state, + step=state.step + 1, + beta=beta, + tripped=tripped, + dwell_count=dwell, + release_count=release, + last_signal=min(1.0, max(0.0, fvote)), + ) + action = MitigationAction( + new_beta=beta, + tripped=tripped, + verdict=_verdict_for(fvote), + reason=reason, + ) + return new_state, action + + class MitigationLogWriter: """Thread-safe append-only JSONL log for the reward-hack controller. @@ -341,6 +482,7 @@ class RewardHackMitigationCallback(_TrainerCallbackBase): # type: ignore[misc, buffer: Any = None, tokenizer: Any = None, task: str = "grpo", + bang_bang: BangBangPolicy | None = None, ) -> None: if mode not in MITIGATION_MODES: raise ValueError( @@ -363,6 +505,9 @@ class RewardHackMitigationCallback(_TrainerCallbackBase): # type: ignore[misc, self.buffer = buffer self.tokenizer = tokenizer self.task = task + self.bang_bang = bang_bang + if mode == "kl_control" and bang_bang is None: + raise ValueError("kl_control mode requires a BangBangPolicy") # Compose the v0.70.0 detector callback for the info_rm/rm_ensemble # baseline + drop_pct logic (DRY — no re-implementation). self._detector_cb = RewardHackCallback( @@ -370,6 +515,7 @@ class RewardHackMitigationCallback(_TrainerCallbackBase): # type: ignore[misc, ) self._trainer: Any = None self._state = ControllerState() + self._length_baseline: float | None = None def attach(self, trainer: Any) -> None: """Store the trainer reference (the β / kl_coef mutation target).""" @@ -386,8 +532,22 @@ class RewardHackMitigationCallback(_TrainerCallbackBase): # type: ignore[misc, return getattr(args, "kl_coef", None) if args is not None else None return getattr(trainer, "beta", None) - def _build_telemetry(self, snapshot: Mapping[str, Any], step: int) -> dict[str, Any]: - """Compute the per-step telemetry entry from a buffer snapshot.""" + def _length_trend(self, length_mean: float) -> float: + """Relative growth of the mean completion length vs its baseline, in + ``[0, 1]``. Rising = the policy is padding output (length hacking).""" + if self._length_baseline is None: + if length_mean > 0: + self._length_baseline = length_mean + return 0.0 + base = self._length_baseline + if base <= 0: + return 0.0 + return min(1.0, max(0.0, (length_mean - base) / base)) + + def _observe( + self, snapshot: Mapping[str, Any], step: int + ) -> tuple[dict[str, Any], dict[str, float]]: + """Compute the telemetry entry + the per-signal vote inputs.""" raw = self._detector_cb.compute_signal(snapshot) drop_pct = 0.0 verdict = "OK" @@ -398,7 +558,15 @@ class RewardHackMitigationCallback(_TrainerCallbackBase): # type: ignore[misc, completions = snapshot.get("completions", []) or [] rewards = snapshot.get("rewards", []) or [] reward_mean, reward_std = reward_mean_std(rewards) - return { + length_mean = mean_token_len(completions) + repetition = mean_repetition(completions) + length_trend = self._length_trend(length_mean) + signals = { + self.detector: drop_pct, + "length_trend": length_trend, + "repetition": repetition, + } + telemetry = { "mode": self.mode, "detector": self.detector, "raw_signal": raw, @@ -407,9 +575,70 @@ class RewardHackMitigationCallback(_TrainerCallbackBase): # type: ignore[misc, "beta": self._current_coefficient(), "reward_mean": reward_mean, "reward_std": reward_std, - "completion_length_mean": mean_token_len(completions), - "repetition": mean_repetition(completions), + "completion_length_mean": length_mean, + "repetition": repetition, + "length_trend": length_trend, } + return telemetry, signals + + def _apply_coefficient(self, value: float) -> None: + """Write the controller's coefficient to the trainer. + + GRPO: β must be dual-written — stock ``GRPOTrainer.compute_loss`` reads + ``self.beta`` (the instance) while Soup's ``_GRPOTrainerVariant`` reads + ``self.args.beta`` (the config). PPO: ``args.kl_coef``. + """ + trainer = self._trainer + if trainer is None: + return + args = getattr(trainer, "args", None) + if self.task == "ppo": + if args is not None: + try: + args.kl_coef = value + except Exception: # noqa: BLE001 — never crash training + pass + return + try: + trainer.beta = value + except Exception: # noqa: BLE001 + pass + if args is not None: + try: + args.beta = value + except Exception: # noqa: BLE001 + pass + + def _seed_coefficient(self, floor: float) -> None: + """Seed the controller β from the live trainer coefficient on step 1.""" + if self._state.beta > 0.0: + return + current = self._current_coefficient() + seed = ( + float(current) + if isinstance(current, (int, float)) + and not isinstance(current, bool) + and current > 0 + else floor + ) + self._state = replace(self._state, beta=seed) + + def _run_bang_bang( + self, telemetry: dict[str, Any], signals: Mapping[str, float] + ) -> None: + """kl_control: vote → bang-bang step → mutate the trainer coefficient.""" + policy = self.bang_bang + if policy is None: + return + self._seed_coefficient(policy.beta_floor) + vote = combine_signals(signals, self.signals) + new_state, action = bang_bang_step(policy, self._state, vote=vote) + self._state = new_state + self._apply_coefficient(action.new_beta) + telemetry["vote"] = vote + telemetry["new_beta"] = action.new_beta + telemetry["tripped"] = action.tripped + telemetry["action"] = action.reason def on_step_end(self, args, state, control, **kwargs): """Per-step hook — read the buffer, compute telemetry, act by mode. @@ -422,10 +651,11 @@ class RewardHackMitigationCallback(_TrainerCallbackBase): # type: ignore[misc, try: snapshot = self.buffer.snapshot() step = int(getattr(state, "global_step", 0) or 0) - telemetry = self._build_telemetry(snapshot, step) - # Stage 0: log_only observes; control modes (kl_control / - # pid_lagrangian) are wired in later stages. In every mode we - # record the telemetry line. + telemetry, signals = self._observe(snapshot, step) + # log_only observes; kl_control drives the bang-bang controller. + # pid_lagrangian is wired in Stage 2. + if self.mode == "kl_control": + self._run_bang_bang(telemetry, signals) self.log_writer.record(step=step, snapshot=telemetry) return control except Exception: # noqa: BLE001 — instrumentation must never crash diff --git a/tests/test_v07126.py b/tests/test_v07126.py index fdc9dde..d7ef616 100644 --- a/tests/test_v07126.py +++ b/tests/test_v07126.py @@ -49,6 +49,24 @@ def _fake_ppo_trainer(kl_coef=0.2): args = types.SimpleNamespace(kl_coef=kl_coef) return types.SimpleNamespace(args=args) + +class _SeqBuffer: + """Fake RLSignalBuffer returning a scripted sequence of snapshots.""" + + def __init__(self, snapshots): + self._snapshots = snapshots + self._index = 0 + + def snapshot(self): + snap = self._snapshots[min(self._index, len(self._snapshots) - 1)] + self._index += 1 + return snap + + +# Well-separated rewards (healthy) vs bunched rewards (reward model losing grip). +_HEALTHY = _grpo_snapshot([0.0, 0.0, 1.0, 1.0], ["a", "b", "c", "d"]) +_HACK = _grpo_snapshot([0.5, 0.5, 0.5, 0.5], ["a", "b", "c", "d"]) + # ===================================================================== # Part A / Stage 0 — schema: reward_hack_mitigation field + gate (Task A1) # ===================================================================== @@ -658,3 +676,364 @@ class TestAttachMitigationWiring: ) assert any(isinstance(c, RewardHackCallback) for c in added) assert not any(isinstance(c, RewardHackMitigationCallback) for c in added) + + +# ===================================================================== +# Part B / Stage 1 — schema fields + validators (Task B1) +# ===================================================================== + + +class TestStage1Schema: + """Bang-bang controller tunables + bounds + mutual-exclusion.""" + + def _cfg(self, extra: str, *, mitigation: str = "kl_control"): + from soup_cli.config.loader import load_config_from_string + + return load_config_from_string(_yaml("grpo", mitigation=mitigation, extra=extra)) + + def test_kl_control_parses(self): + cfg = self._cfg("reward_hack_beta_floor: 0.05\nreward_hack_beta_ceil: 2.0") + assert cfg.training.reward_hack_mitigation == "kl_control" + assert cfg.training.reward_hack_beta_floor == 0.05 + assert cfg.training.reward_hack_beta_ceil == 2.0 + + def test_defaults(self): + from soup_cli.config.schema import TrainingConfig + + t = TrainingConfig() + assert t.reward_hack_beta_floor == 0.02 and t.reward_hack_beta_ceil == 1.0 + assert t.reward_hack_trip_band == 0.30 and t.reward_hack_release_band == 0.10 + assert t.reward_hack_dwell_steps == 2 and t.reward_hack_release_patience == 3 + assert t.reward_hack_kl_gain == 1.5 and t.reward_hack_signals == ["info_rm"] + + def test_floor_ge_ceil_rejected(self): + with pytest.raises(ValueError, match="beta_floor"): + self._cfg("reward_hack_beta_floor: 1.0\nreward_hack_beta_ceil: 0.5") + + def test_release_ge_trip_rejected(self): + with pytest.raises(ValueError, match="release_band"): + self._cfg("reward_hack_trip_band: 0.2\nreward_hack_release_band: 0.3") + + def test_unknown_signal_rejected(self): + with pytest.raises(ValueError, match="signal"): + self._cfg("reward_hack_signals: [info_rm, not_a_signal]") + + def test_known_signals_accepted(self): + cfg = self._cfg("reward_hack_signals: [info_rm, length_trend, repetition]") + assert cfg.training.reward_hack_signals == ["info_rm", "length_trend", "repetition"] + + def test_kl_control_excludes_ref_ema(self): + with pytest.raises(ValueError, match="ref_model_ema_alpha"): + self._cfg("ref_model_ema_alpha: 0.9") + + def test_beta_floor_must_be_positive(self): + with pytest.raises(ValueError): + self._cfg("reward_hack_beta_floor: 0.0") + + def test_kl_gain_must_exceed_one(self): + with pytest.raises(ValueError): + self._cfg("reward_hack_kl_gain: 1.0") + + def test_dwell_must_be_positive(self): + with pytest.raises(ValueError): + self._cfg("reward_hack_dwell_steps: 0") + + def test_tunable_without_mode_rejected(self): + # footgun: setting a control tunable while mitigation is off. + from soup_cli.config.loader import load_config_from_string + + with pytest.raises(ValueError, match="reward_hack_mitigation"): + load_config_from_string( + _yaml( + "grpo", detector=None, mitigation="off", extra="reward_hack_kl_gain: 2.0" + ) + ) + + +# ===================================================================== +# Part B / Stage 1 — BangBangPolicy + bang_bang_step (Task B2) +# ===================================================================== + + +def _bang_policy(**kw): + from soup_cli.utils.reward_hack_control import BangBangPolicy + + defaults = dict( + beta_floor=0.02, + beta_ceil=1.0, + trip_band=0.3, + release_band=0.1, + dwell_steps=2, + release_patience=2, + kl_gain=1.5, + ) + defaults.update(kw) + return BangBangPolicy(**defaults) + + +def _run_bang(policy, votes, beta0=None): + from soup_cli.utils.reward_hack_control import ControllerState, bang_bang_step + + state = ControllerState(beta=beta0 if beta0 is not None else policy.beta_floor) + actions = [] + for vote in votes: + state, action = bang_bang_step(policy, state, vote=vote) + actions.append(action) + return state, actions + + +class TestBangBang: + """Reversible bang-bang controller with dwell + release hysteresis.""" + + def test_policy_floor_lt_ceil(self): + from soup_cli.utils.reward_hack_control import BangBangPolicy + + with pytest.raises(ValueError, match="floor"): + BangBangPolicy( + beta_floor=1.0, + beta_ceil=0.5, + trip_band=0.3, + release_band=0.1, + dwell_steps=2, + release_patience=2, + kl_gain=1.5, + ) + + def test_policy_release_lt_trip(self): + with pytest.raises(ValueError, match="release"): + _bang_policy(trip_band=0.2, release_band=0.3) + + def test_policy_kl_gain_gt_one(self): + with pytest.raises(ValueError, match="kl_gain"): + _bang_policy(kl_gain=1.0) + + def test_policy_beta_floor_positive(self): + with pytest.raises(ValueError, match="beta_floor"): + _bang_policy(beta_floor=0.0) + + def test_below_band_no_trip(self): + state, _ = _run_bang(_bang_policy(), [0.1, 0.1, 0.1, 0.1]) + assert not state.tripped and state.beta == pytest.approx(0.02) + + def test_dwell_then_trip(self): + state, _ = _run_bang(_bang_policy(), [0.5, 0.5]) + assert state.tripped and state.beta == pytest.approx(0.03) # 0.02 * 1.5 + + def test_keeps_raising_while_hacking(self): + state, _ = _run_bang(_bang_policy(), [0.5, 0.5, 0.5, 0.5]) + assert state.tripped and state.beta == pytest.approx(0.02 * 1.5**3) + + def test_no_flap_on_alternating(self): + # A naive controller would flap; dwell + release hysteresis prevents it. + state, _ = _run_bang(_bang_policy(), [0.5, 0.05, 0.5, 0.05, 0.5, 0.05]) + assert not state.tripped and state.beta == pytest.approx(0.02) + + def test_release_reverses(self): + # trip (2 steps) then release_patience (2 steps low) → β back to floor. + state, _ = _run_bang(_bang_policy(), [0.5, 0.5, 0.05, 0.05]) + assert state.beta == pytest.approx(0.02) and not state.tripped + + def test_beta_clamped_to_ceil(self): + state, _ = _run_bang(_bang_policy(beta_ceil=0.05), [0.5] * 10) + assert state.beta == pytest.approx(0.05) + + def test_action_is_frozen(self): + from dataclasses import FrozenInstanceError + + from soup_cli.utils.reward_hack_control import ControllerState, bang_bang_step + + _, action = bang_bang_step(_bang_policy(), ControllerState(beta=0.02), vote=0.5) + with pytest.raises(FrozenInstanceError): + action.new_beta = 1.0 # type: ignore[misc] + + def test_action_verdict_and_reason(self): + from soup_cli.utils.reward_hack_control import ControllerState, bang_bang_step + + _, action = bang_bang_step( + _bang_policy(dwell_steps=1), ControllerState(beta=0.02), vote=0.9 + ) + assert action.verdict == "HACK" and "raise" in action.reason.lower() + + +# ===================================================================== +# Part B / Stage 1 — kl_control callback: β dual-write + kl_coef (Task B3) +# ===================================================================== + + +def _kl_policy(**kw): + from soup_cli.utils.reward_hack_control import BangBangPolicy + + defaults = dict( + beta_floor=0.02, + beta_ceil=1.0, + trip_band=0.3, + release_band=0.1, + dwell_steps=1, + release_patience=1, + kl_gain=2.0, + ) + defaults.update(kw) + return BangBangPolicy(**defaults) + + +def _kl_callback(tmp_path, buffer, *, task="grpo", signals=("info_rm",), policy=None): + from soup_cli.utils.reward_hack_control import ( + MitigationLogWriter, + RewardHackMitigationCallback, + ) + + return RewardHackMitigationCallback( + mode="kl_control", + detector="info_rm", + log_writer=MitigationLogWriter(str(tmp_path / "m.jsonl")), + signals=signals, + buffer=buffer, + task=task, + bang_bang=policy or _kl_policy(), + ) + + +class TestKlControlCallback: + """kl_control mutates β (GRPO, dual-write) / kl_coef (PPO) via the controller.""" + + def test_kl_control_requires_policy(self, tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + from soup_cli.utils.reward_hack_control import ( + MitigationLogWriter, + RewardHackMitigationCallback, + ) + + with pytest.raises(ValueError, match="kl_control"): + RewardHackMitigationCallback( + mode="kl_control", + detector="info_rm", + log_writer=MitigationLogWriter("m.jsonl"), + ) + + def test_grpo_dual_write_on_hack(self, tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + cb = _kl_callback(tmp_path, _SeqBuffer([_HEALTHY, _HACK])) + trainer = _fake_grpo_trainer(beta=0.02) + cb.attach(trainer) + cb.on_step_end(None, types.SimpleNamespace(global_step=1), None) # baseline + cb.on_step_end(None, types.SimpleNamespace(global_step=2), None) # hack → raise + assert trainer.beta == pytest.approx(0.04) # 0.02 * 2.0 + assert trainer.args.beta == pytest.approx(0.04) # DUAL write (variant path) + + def test_ppo_kl_coef_mutated(self, tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + cb = _kl_callback(tmp_path, _SeqBuffer([_HEALTHY, _HACK]), task="ppo") + trainer = _fake_ppo_trainer(kl_coef=0.2) + cb.attach(trainer) + cb.on_step_end(None, types.SimpleNamespace(global_step=1), None) + cb.on_step_end(None, types.SimpleNamespace(global_step=2), None) + assert trainer.args.kl_coef == pytest.approx(0.4) # 0.2 * 2.0 + + def test_recovery_relaxes_beta(self, tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + cb = _kl_callback(tmp_path, _SeqBuffer([_HEALTHY, _HACK, _HEALTHY])) + trainer = _fake_grpo_trainer(beta=0.02) + cb.attach(trainer) + for step in (1, 2, 3): # baseline, hack (raise), recovery (relax) + cb.on_step_end(None, types.SimpleNamespace(global_step=step), None) + assert trainer.beta == pytest.approx(0.02) # relaxed back to floor + assert not cb._state.tripped + + def test_kl_control_logs_vote_and_action(self, tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + cb = _kl_callback(tmp_path, _SeqBuffer([_HEALTHY, _HACK])) + cb.attach(_fake_grpo_trainer(beta=0.02)) + cb.on_step_end(None, types.SimpleNamespace(global_step=1), None) + cb.on_step_end(None, types.SimpleNamespace(global_step=2), None) + lines = (tmp_path / "m.jsonl").read_text().strip().splitlines() + last = json.loads(lines[-1]) + for key in ("vote", "new_beta", "tripped", "action"): + assert key in last + + +class TestAttachKlControl: + """attach_rl_callbacks builds the bang-bang policy from tcfg for kl_control.""" + + def test_attach_kl_control_builds_policy(self, tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + from soup_cli.config.schema import TrainingConfig + from soup_cli.utils.peft_wiring import attach_rl_callbacks + from soup_cli.utils.reward_hack_control import RewardHackMitigationCallback + + tcfg = TrainingConfig( + reward_hack_mitigation="kl_control", + reward_hack_detector="info_rm", + reward_hack_kl_gain=3.0, + reward_hack_signals=["info_rm", "length_trend"], + ) + added: list = [] + attach_rl_callbacks( + _fake_trainer_recording(added), + tcfg, + buffer=object(), + output_dir=str(tmp_path), + task="grpo", + ) + mit = [c for c in added if isinstance(c, RewardHackMitigationCallback)] + assert len(mit) == 1 + assert mit[0].mode == "kl_control" + assert mit[0].bang_bang is not None and mit[0].bang_bang.kl_gain == 3.0 + assert mit[0].signals == ("info_rm", "length_trend") + + +class TestPpoBufferParity: + """PPO wires the RLSignalBuffer + mitigation callback (BETA — GRPO gets the + on-GPU proof; PPO's kl_coef mutation is unit-tested in TestKlControlCallback).""" + + def test_ppo_setup_wires_buffer_and_mitigation(self): + import inspect + + from soup_cli.trainer import ppo + + src = inspect.getsource(ppo) + assert "RLSignalBuffer" in src + assert "rl_callbacks_need_buffer" in src + assert "attach_rl_callbacks" in src + assert 'task="ppo"' in src or "task='ppo'" in src + + +class TestRewardHackMitigationCli: + """`soup train --reward-hack-mitigation ` flag + re-exec passthrough.""" + + def _runner(self): + from typer.testing import CliRunner + + from soup_cli.cli import app + + return CliRunner(), app + + def test_help_shows_flag(self): + runner, app = self._runner() + result = runner.invoke(app, ["train", "--help"]) + assert result.exit_code == 0, result.output + assert "--reward-hack-mitigation" in result.output + + def test_override_without_detector_rejected(self, tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + (tmp_path / "cfg.yaml").write_text( + "base: HuggingFaceTB/SmolLM2-135M\n" + "task: grpo\n" + "data:\n train: ./train.jsonl\n format: chatml\n" + "training:\n reward_fn: accuracy\n" + ) + runner, app = self._runner() + result = runner.invoke( + app, + ["train", "--config", "cfg.yaml", "--reward-hack-mitigation", "log_only", "--yes"], + ) + assert result.exit_code == 1 + assert "reward_hack_detector" in result.output + + def test_reexec_passthrough_present(self): + import inspect + + from soup_cli.commands import train as train_mod + + src = inspect.getsource(train_mod) + assert "--reward-hack-mitigation" in src + assert "cfg.training.reward_hack_mitigation" in src