Codeseys/composer-replication-framework
0
1"""composer_trainer.py — TRL GRPOTrainer subclass with SDPO + trace-replay channels.2 3Architecture spec: docs/INTEGRATION_ARCHITECTURE.md § "Recipe A".4Verified extension point: GRPOTrainer._compute_loss(model, inputs)5 (DeepWiki audit of huggingface/trl, 2026-05-25).6 7Total loss:8 total_loss = grpo_loss9 + alpha_sdpo * sdpo_kl_at_error_turns10 + beta_replay * trace_replay_dpo_loss11 12Where:13 - grpo_loss is the parent GRPOTrainer's loss (RLVR + DAPO patches).14 - sdpo_kl_at_error_turns is generalized_jsd_loss between student's logits and15 teacher's (= same-model-with-hint-context) logits, masked to error-turn tokens only.16 - trace_replay_dpo_loss is DPO loss over (chosen, rejected) pairs derived from17 N external teacher disagreement with the student.18 19The data collator (data_collator.py) is responsible for:20 - Detecting error sites in the rollout and constructing ctx_teacher = ctx_student + hint.21 - Computing sdpo_loss_mask (1 at post-hint error-turn tokens, 0 elsewhere).22 - Loading DPO pairs from the trace-replay output (see teacher_replay.py).23 - Precomputing reference-policy logprobs for DPO.24"""25 26from __future__ import annotations27 28import logging29from collections.abc import Callable30from typing import TYPE_CHECKING, Any31 32import torch33import torch.nn.functional as F # noqa: N812 — repo-wide torch convention34 35if TYPE_CHECKING: # type-only — never imported at runtime (keeps the dep lazy)36 from composer_replication.safety import HeldOutGuard37 38# These imports work when TRL is installed — they're not skeleton imports.39# When TRL is missing we fall back to `object` so the module still imports40# (e.g. for documentation generation) but raise a clear ImportError at41# instantiation time rather than the cryptic `object.__init__()` error.42try:43 from trl import GRPOTrainer # type: ignore44 _TRL_AVAILABLE = True45except ImportError: # pragma: no cover — only hit in unit-test stubs without TRL46 GRPOTrainer = object # type: ignore — fallback so module imports without TRL47 _TRL_AVAILABLE = False48 49from composer_replication.opsd import generalized_jsd_loss50from composer_replication.trainer.kl_in_reward import (51 apply_kl_in_reward,52 kl_penalty_per_sequence,53)54 55logger = logging.getLogger(__name__)56 57 58class ComposerReplicationTrainer(GRPOTrainer): # type: ignore[misc, valid-type]59 """TRL GRPOTrainer with Composer-recipe channels (SDPO) + novel trace-replay-DPO.60 61 Args (in addition to GRPOTrainer's):62 alpha_sdpo: weight on SDPO hint-distill loss. Default 0.0 (disabled).63 Opt in by passing >0 once your data collator produces64 `sdpo_loss_mask` and `ctx_teacher_input_ids` columns.65 beta_replay: weight on trace-replay DPO loss. Default 0.0 (disabled).66 Opt in by passing >0 once your data collator produces67 `dpo_chosen_input_ids` / `dpo_rejected_input_ids` etc.68 sdpo_jsd_beta: beta param of generalized_jsd_loss69 (0=KL(teacher||student), 0.5=JSD, 1=KL(student||teacher) per70 upstream OPSD convention; see composer_replication/opsd.py).71 sdpo_temperature: temperature for SDPO loss; SDPO paper uses 1.0.72 sdpo_token_clip: per-token JSD clip for stability; None = no clip.73 replay_dpo_beta: beta param of the DPO loss (β in the standard DPO formula).74 kl_in_reward: when True, apply the KL-to-reference penalty in the75 **reward** (Composer-2 §4.1 / verl choice) instead of TRL's native76 **in-loss** k3 term. The penalty is folded into GRPO's advantages at77 scoring time (``adv -= beta·(KL - group_mean(KL))``) and TRL's78 in-loss KL is suppressed for that step. The F5 audit's #1 fidelity79 fix: the 2025/26 evidence (arXiv:2512.21852, verl, TRL #4967) shows80 k1-in-reward improves OOD generalization where k3-in-reward can81 collapse. REQUIRES ``beta>0`` (the KL coefficient — also how TRL82 decides to compute reference logprobs) and ``scale_rewards`` in83 {none,false} (the advantage-adjustment identity is exact only84 without std-normalization — the Dr.GRPO / Composer regime). Default85 False = TRL's native in-loss KL, byte-for-byte legacy behavior.86 kl_estimator: ``"k1"`` (default; ``logp - ref_logp``, the Composer-2 /87 verl choice this path exists for) or ``"k3"`` (Schulman; lets an88 experiment A/B k1-in-reward vs k3-in-reward). Only consulted when89 ``kl_in_reward=True``.90 heldout_guard: optional ``HeldOutGuard`` (the #2 collapse safeguard from91 ``composer_replication.safety``). Default None = OFF (no behavior92 change whatsoever). When supplied, the trainer folds one checkpoint's93 metrics into the guard at the ``args.logging_steps`` cadence (the same94 place the loss components are logged) and HALTS the run on a fired95 verdict — the run-level reward-hacking / collapse tripwire actually96 firing instead of sitting inert.97 heldout_eval_fn: zero-arg callable returning the held-out (real) eval98 score as a float, evaluated each guard cadence. Injectable so the99 trainer never hardcodes an eval — pass a closure over your disjoint100 held-out pool (the ``HeldoutSplit`` discipline). Required whenever101 ``heldout_guard`` is set; the guard's whole signal is in-loop reward102 vs. this held-out score.103 strict_killswitch: when True (default), a fired guard verdict raises104 ``CollapseStopError`` to hard-stop training (exception-based control105 flow, matching ``HeldOutGuard.raise_if_fired``). When False the106 verdict is logged and ``self.control.should_training_stop`` is set so107 the HF loop ends gracefully after the step (soft stop). Only consulted108 when ``heldout_guard`` is set.109 """110 111 def __init__(112 self,113 *args: Any,114 alpha_sdpo: float = 0.0,115 beta_replay: float = 0.0,116 sdpo_jsd_beta: float = 0.5,117 sdpo_temperature: float = 1.0,118 sdpo_token_clip: float | None = None,119 replay_dpo_beta: float = 0.1,120 strict_sdpo_alignment: bool = True,121 kl_in_reward: bool = False,122 kl_estimator: str = "k1",123 heldout_guard: HeldOutGuard | None = None,124 heldout_eval_fn: Callable[[], float] | None = None,125 strict_killswitch: bool = True,126 **kwargs: Any,127 ):128 if not _TRL_AVAILABLE:129 raise ImportError(130 "ComposerReplicationTrainer requires TRL. Install with "131 "`pip install -e .[train]`."132 )133 super().__init__(*args, **kwargs)134 self.alpha_sdpo = alpha_sdpo135 self.beta_replay = beta_replay136 self.sdpo_jsd_beta = sdpo_jsd_beta137 self.sdpo_temperature = sdpo_temperature138 self.sdpo_token_clip = sdpo_token_clip139 self.replay_dpo_beta = replay_dpo_beta140 # When True (default), an SDPO student/teacher shape mismatch is a hard141 # error — it means the data collator failed to align the post-hint142 # section, which silently zeroes the distillation signal (the exact143 # trust-gap flagged in ADR-008). Set False only for production runs144 # where a single malformed batch should warn-and-skip rather than abort.145 self.strict_sdpo_alignment = strict_sdpo_alignment146 # --- k1-in-reward KL (F5 #1 fidelity fix; Composer-2 §4.1 / verl) ----147 # OFF by default → TRL's native in-loss k3 KL, byte-for-byte legacy.148 # When ON we keep self.beta as the KL coef (TRL needs beta>0 to even149 # create the ref model + compute ref logps), fold the k1 penalty into150 # advantages during scoring, and zero TRL's in-loss KL per step.151 self.kl_in_reward = kl_in_reward152 self.kl_estimator = kl_estimator153 if kl_in_reward:154 validate_kl_in_reward_config(155 kl_estimator=kl_estimator,156 beta=float(getattr(self.args, "beta", 0.0)),157 scale_rewards=getattr(self.args, "scale_rewards", "group"),158 )159 # --- run-level collapse kill-switch (#2 safeguard) -------------------160 # OPTIONAL + OFF BY DEFAULT: when heldout_guard is None the loss path is161 # byte-for-byte the legacy behavior. When set, _maybe_update_killswitch162 # folds metrics into the guard at the logging cadence (see _compute_loss).163 self.heldout_guard = heldout_guard164 self.heldout_eval_fn = heldout_eval_fn165 self.strict_killswitch = strict_killswitch166 if heldout_guard is not None and heldout_eval_fn is None:167 raise ValueError(168 "heldout_guard was provided without heldout_eval_fn: the guard's "169 "tripwire compares in-loop reward against a DISJOINT held-out "170 "(real) eval score, so it needs an injectable zero-arg "171 "heldout_eval_fn() -> float. Pass a closure over your held-out "172 "pool (the HeldoutSplit discipline)."173 )174 175 # ----------------------------------------------------------------------176 # Loss override (the integration core)177 # ----------------------------------------------------------------------178 179 # ----------------------------------------------------------------------180 # k1-in-reward: fold the KL penalty into advantages at scoring time, and181 # suppress TRL's native in-loss k3 KL inside _compute_loss.182 # ----------------------------------------------------------------------183 184 def _generate_and_score_completions(185 self,186 inputs: list[dict[str, Any]],187 ) -> dict[str, Any]:188 """Override: after TRL scores completions, fold a k1 KL penalty into the189 advantages (Composer-2 in-reward KL) when ``kl_in_reward`` is set.190 191 No-op (exact legacy path) when ``kl_in_reward`` is False. When set, TRL192 has already computed ``advantages``, ``ref_per_token_logps`` (because193 ``beta>0``), and the completion logprobs; we recompute the per-sequence194 k1 penalty and apply the exact group-mean-baseline correction.195 """196 output = super()._generate_and_score_completions(inputs)197 if not getattr(self, "kl_in_reward", False):198 return output199 200 ref_logps = output.get("ref_per_token_logps")201 # The "old" (sampling-time) policy logps are TRL's in-loss π term; they202 # may be lazily None when generation/optimization are aligned and not203 # vLLM (see TRL _compute_loss: old := per_token_logps.detach()). In that204 # aligned case we cannot read π logps here, so we defer to _compute_loss205 # (which always has per_token_logps) by stashing what we need.206 old_logps = output.get("old_per_token_logps")207 completion_mask = output.get("completion_mask")208 if ref_logps is None or completion_mask is None:209 # beta>0 guarantees ref_logps; this branch only trips on a TRL210 # internals change — fail loud rather than silently skip the penalty.211 raise RuntimeError(212 "kl_in_reward=True but TRL did not return ref_per_token_logps / "213 "completion_mask from scoring (beta>0 should guarantee them). "214 "TRL internals may have changed; re-verify the in-reward path."215 )216 217 if old_logps is not None:218 penalty = kl_penalty_per_sequence(219 policy_logps=old_logps,220 ref_logps=ref_logps,221 completion_mask=completion_mask,222 estimator=self.kl_estimator,223 )224 output["advantages"] = apply_kl_in_reward(225 advantages=output["advantages"],226 kl_penalty=penalty,227 num_generations=self.num_generations,228 coef=float(self.args.beta),229 )230 output["_kl_in_reward_applied"] = torch.tensor(True)231 else:232 # Aligned non-vLLM case: π logps materialize only in _compute_loss.233 # Stash ref logps + mask so _compute_loss can apply the penalty there.234 output["_kl_in_reward_applied"] = torch.tensor(False)235 return output236 237 def _compute_loss(238 self,239 model: torch.nn.Module,240 inputs: dict[str, torch.Tensor],241 ) -> torch.Tensor:242 """Override: total_loss = grpo + α*sdpo + β*replay.243 244 When ``kl_in_reward`` is set, TRL's native in-loss KL term (gated on245 ``self.beta``) is suppressed by temporarily zeroing ``self.beta`` for the246 duration of the parent call — the KL has already been (or is about to be)247 accounted for in the reward/advantage, so double-counting it in the loss248 would be wrong. ``self.beta`` is restored in ``finally``.249 """250 # Channel 1: standard GRPO loss. ``getattr`` (not ``self.kl_in_reward``)251 # so an instance built via ``__new__`` + manual wiring (the SDPO /252 # kill-switch unit-test pattern that skips __init__) defaults to the253 # legacy path instead of raising AttributeError.254 if getattr(self, "kl_in_reward", False):255 grpo_loss = self._grpo_loss_kl_in_reward(model, inputs)256 else:257 grpo_loss = super()._compute_loss(model, inputs)258 259 # Channel 2: SDPO hint-distill at error sites260 sdpo_kl = self._compute_sdpo_loss(model, inputs)261 262 # Channel 3: trace-replay DPO from teacher disagreement263 replay_dpo = self._compute_trace_replay_loss(model, inputs)264 265 # Compose266 total = grpo_loss + self.alpha_sdpo * sdpo_kl + self.beta_replay * replay_dpo267 268 # Log per-channel components (so we can ablate post-hoc)269 if hasattr(self, "state") and getattr(self, "args", None) is not None:270 log_steps = getattr(self.args, "logging_steps", 50)271 if self.state.global_step % log_steps == 0:272 self.log({ # type: ignore[attr-defined]273 "loss/grpo": float(grpo_loss.detach()),274 "loss/sdpo_kl": float(sdpo_kl.detach()),275 "loss/trace_replay_dpo": float(replay_dpo.detach()),276 "loss/total": float(total.detach()),277 "loss/alpha_sdpo": self.alpha_sdpo,278 "loss/beta_replay": self.beta_replay,279 })280 # Fold one checkpoint into the run-level collapse kill-switch at281 # the SAME cadence (no-op unless a guard was configured).282 self._maybe_update_killswitch()283 284 return total285 286 def _grpo_loss_kl_in_reward(287 self,288 model: torch.nn.Module,289 inputs: dict[str, torch.Tensor],290 ) -> torch.Tensor:291 """GRPO loss with the KL applied in the reward, not the loss.292 293 Two responsibilities:294 1. Suppress TRL's native in-loss k3 KL term for this step by zeroing295 ``self.beta`` across the parent ``_compute_loss`` call (restored in296 ``finally``). ``self.beta`` gates the in-loss KL add (TRL297 ``_compute_loss``: ``if self.beta != 0.0: per_token_loss += beta*kl``).298 2. Handle the deferred case: when generation/optimization are aligned299 and not using vLLM, the sampling-time policy logps are None at300 scoring time, so ``_generate_and_score_completions`` could not fold301 the penalty into advantages. Here ``per_token_logps`` is available,302 so we apply the exact same advantage correction in-place on303 ``inputs["advantages"]`` BEFORE the parent computes the surrogate.304 """305 # Deferred-penalty path: advantages not yet KL-adjusted (aligned, no vLLM).306 applied = inputs.get("_kl_in_reward_applied")307 already_applied = bool(applied.item()) if applied is not None else False308 if not already_applied and "ref_per_token_logps" in inputs:309 with torch.no_grad():310 prompt_ids, completion_ids = inputs["prompt_ids"], inputs["completion_ids"]311 completion_mask = inputs["completion_mask"]312 input_ids = torch.cat([prompt_ids, completion_ids], dim=1)313 attention_mask = torch.cat([inputs["prompt_mask"], completion_mask], dim=1)314 logits_to_keep = completion_ids.size(1)315 policy_logps, _ = self._get_per_token_logps_and_entropies(316 model, input_ids, attention_mask, logits_to_keep317 )318 penalty = kl_penalty_per_sequence(319 policy_logps=policy_logps,320 ref_logps=inputs["ref_per_token_logps"],321 completion_mask=completion_mask,322 estimator=self.kl_estimator,323 )324 advantages = inputs["advantages"]325 # advantages may be (B,) or (B,1) — squeeze for the penalty math,326 # restore the original shape after.327 adv_flat = advantages.reshape(advantages.shape[0])328 adj = apply_kl_in_reward(329 advantages=adv_flat,330 kl_penalty=penalty,331 num_generations=self.num_generations,332 coef=float(self.args.beta),333 )334 inputs["advantages"] = adj.reshape(advantages.shape)335 336 # Suppress TRL's in-loss KL: zero beta for the parent call, restore after.337 saved_beta = self.beta338 try:339 self.beta = 0.0340 return super()._compute_loss(model, inputs)341 finally:342 self.beta = saved_beta343 344 # ----------------------------------------------------------------------345 # Run-level collapse kill-switch (#2 safeguard) — optional, OFF by default346 # ----------------------------------------------------------------------347 348 def _maybe_update_killswitch(self) -> None:349 """Fold this checkpoint's metrics into ``heldout_guard`` and act on a fire.350 351 No-op when no guard was configured (the default) — this is the352 backward-compat guarantee: without ``heldout_guard`` the trainer behaves353 exactly as before. When a guard IS set:354 355 * ``in_loop_reward`` is the GRPO reward signal TRL already aggregates356 into ``self._metrics[mode]["reward"]`` each step (we read the latest;357 no extra forward pass).358 * ``heldout_score`` comes from the injected ``heldout_eval_fn()`` — the359 trainer never hardcodes an eval.360 * ``kl_to_init`` (token-mean nats/token, the ``token_mean_kl``361 convention the guard expects) is read from TRL's logged ``"kl"``362 metric when present, else left None (KL path stays inert).363 364 On a fired verdict the verdict is logged. If ``strict_killswitch`` (the365 default) the verdict is converted into a ``CollapseStopError`` via366 ``HeldOutGuard.raise_if_fired`` (hard stop); otherwise the HF training367 loop is asked to stop gracefully after this step.368 """369 guard = self.heldout_guard370 if guard is None:371 return # OFF by default — zero behavior change372 373 round_idx = int(getattr(self.state, "global_step", 0))374 in_loop_reward = self._latest_metric("reward")375 if in_loop_reward is None:376 # No reward aggregated yet (e.g. very first micro-step before TRL has377 # populated its metrics). Skip this cadence rather than feed a378 # fabricated 0.0 that would pollute the guard's baseline/EMA.379 logger.debug(380 "kill-switch: no in-loop reward metric yet at step %d; skipping.",381 round_idx,382 )383 return384 385 assert self.heldout_eval_fn is not None # enforced in __init__386 heldout_score = float(self.heldout_eval_fn())387 kl_to_init = self._latest_metric("kl") # token-mean KL, or None388 389 status = guard.update(390 round_idx=round_idx,391 in_loop_reward=in_loop_reward,392 heldout_score=heldout_score,393 kl_to_init=kl_to_init,394 )395 396 self.log({ # type: ignore[attr-defined]397 "killswitch/in_loop_reward": status.in_loop_ema,398 "killswitch/heldout_score": status.heldout_ema,399 "killswitch/proxy_real_gap": status.proxy_real_gap,400 "killswitch/fire": float(status.fire),401 })402 403 if status.fire:404 logger.error(405 "HeldOutGuard FIRED at step %d — halting run. reason: %s",406 round_idx, status.reason,407 )408 if self.strict_killswitch:409 # Typed exception — exception-based hard stop.410 guard.raise_if_fired(status)411 else:412 # Soft stop: let the HF loop terminate gracefully after this step.413 control = getattr(self, "control", None)414 if control is not None:415 control.should_training_stop = True416 417 def _latest_metric(self, name: str) -> float | None:418 """Most-recent value of a TRL-aggregated train metric, or None.419 420 TRL's GRPOTrainer appends per-step aggregates to421 ``self._metrics["train"][name]`` (e.g. ``"reward"``, ``"kl"``). We read422 the tail defensively so a TRL internals rename degrades to None (KL/reward423 path goes inert) rather than crashing training.424 """425 metrics = getattr(self, "_metrics", None)426 if not isinstance(metrics, dict):427 return None428 train = metrics.get("train")429 if not isinstance(train, dict):430 return None431 series = train.get(name)432 if not series:433 return None434 try:435 return float(series[-1])436 except (TypeError, ValueError, IndexError):437 return None438 439 # ----------------------------------------------------------------------440 # Channel 2: SDPO hint-distill441 # ----------------------------------------------------------------------442 443 def _compute_sdpo_loss(444 self,445 model: torch.nn.Module,446 inputs: dict[str, torch.Tensor],447 ) -> torch.Tensor:448 """Compute generalized_jsd_loss between student and hint-conditioned teacher.449 450 Both come from the SAME model — teacher just has hint inserted into context.451 Skipped (returns 0) if the batch has no error sites (data collator emits452 empty ctx_teacher_input_ids).453 """454 if (455 self.alpha_sdpo == 0.0456 or "ctx_teacher_input_ids" not in inputs457 or inputs["ctx_teacher_input_ids"].numel() == 0458 ):459 return torch.tensor(0.0, device=_device_of(model), requires_grad=True)460 461 # Student forward (with grad, on the original-context input)462 student_logits = model(input_ids=inputs["input_ids"]).logits463 464 # Teacher forward (no grad — same model, hint-conditioned context)465 with torch.no_grad():466 teacher_logits = model(input_ids=inputs["ctx_teacher_input_ids"]).logits467 468 # ------------------------------------------------------------------469 # ALIGNMENT (cross-family review 2026-05-29 — the 4/4-reviewer P0).470 #471 # The teacher context has a hint inserted at the error turn, so the472 # teacher's post-hint response tokens are shifted right by len(hint)473 # relative to the student's. A bare `student.shape == teacher.shape`474 # check does NOT establish token-level alignment: equal-length tensors475 # whose response regions are offset will be JSD'd position-by-position476 # against each other, distilling garbage into the policy.477 #478 # The ONLY correct alignment is an explicit map from the collator that479 # selects, for each response token, the matching index in each sequence.480 # We require it whenever SDPO is active:481 # - `student_response_idx` / `teacher_response_idx`: LongTensors of482 # equal length selecting the aligned response positions in each483 # sequence (the collator builds these knowing where it inserted the484 # hint). JSD is computed over the gathered, provably-aligned logits.485 # - If the collator cannot yet supply them, strict mode raises (loud486 # failure) rather than silently distilling misaligned tokens.487 s_idx = inputs.get("student_response_idx")488 t_idx = inputs.get("teacher_response_idx")489 if s_idx is None or t_idx is None:490 msg = (491 "SDPO alignment indices missing: the collator must emit "492 "`student_response_idx` and `teacher_response_idx` (matching "493 "LongTensors selecting the aligned post-hint response tokens) so "494 "the JSD compares corresponding tokens. A shape-only check does "495 "NOT establish alignment — the hint shifts the teacher's response "496 "tokens right, so equal-length sequences can still be misaligned "497 "and silently distill garbage into the policy (ADR-008 trust-gap)."498 )499 if self.strict_sdpo_alignment:500 raise ValueError(501 msg + " (strict_sdpo_alignment=True; pass False to fall back "502 "to the legacy shape-only check for resilience.)"503 )504 logger.warning("%s Falling back to shape-only alignment check.", msg)505 if student_logits.shape != teacher_logits.shape:506 logger.warning(507 "SDPO shape mismatch student=%s teacher=%s; skipping.",508 tuple(student_logits.shape), tuple(teacher_logits.shape),509 )510 return torch.tensor(0.0, device=_device_of(model), requires_grad=True)511 return generalized_jsd_loss(512 student_logits=student_logits,513 teacher_logits=teacher_logits,514 labels=inputs.get("sdpo_loss_mask"),515 beta=self.sdpo_jsd_beta,516 temperature=self.sdpo_temperature,517 token_clip=self.sdpo_token_clip,518 reduction="batchmean",519 )520 521 # Validate the index tensors describe a consistent 1:1 alignment.522 if s_idx.shape != t_idx.shape:523 raise ValueError(524 f"SDPO alignment index shape mismatch: student_response_idx="525 f"{tuple(s_idx.shape)} vs teacher_response_idx={tuple(t_idx.shape)}. "526 "They must select the same number of aligned response tokens."527 )528 # Gather the provably-aligned response logits from each sequence, then529 # JSD only those positions (this is the masked error-turn distillation).530 # gather over the sequence dim (dim=1): expand index to the vocab dim.531 #532 # ADR-011: ragged-K rows are padded with a sentinel (-1) and a per-row533 # *_valid mask. Negative indices are illegal for torch.gather, so clamp534 # to 0 before gathering, then neutralize those positions by feeding535 # labels=-100 (the standard HF ignore convention that generalized_jsd_loss536 # already honors). This makes sentinel/padding positions contribute 0.537 #538 # Final-verify 2026-05-29: combine BOTH valid masks (not just student's)539 # AND the sentinel guard. If a future collator ever emits divergent540 # student/teacher valid tails, a teacher sentinel clamped to 0 would541 # otherwise be silently distilled against teacher position 0. Belt-and-542 # suspenders: valid iff student-valid AND teacher-valid AND both indices543 # non-sentinel.544 s_valid = inputs.get("student_response_valid")545 t_valid = inputs.get("teacher_response_valid")546 aligned_mask = (s_idx >= 0) & (t_idx >= 0)547 if s_valid is not None:548 aligned_mask = aligned_mask & s_valid.bool()549 if t_valid is not None:550 aligned_mask = aligned_mask & t_valid.bool()551 552 vocab = student_logits.size(-1)553 s_safe = s_idx.clamp_min(0)554 t_safe = t_idx.clamp_min(0)555 s_gather = s_safe.unsqueeze(-1).expand(-1, -1, vocab)556 t_gather = t_safe.unsqueeze(-1).expand(-1, -1, vocab)557 student_aligned = torch.gather(student_logits, 1, s_gather)558 teacher_aligned = torch.gather(teacher_logits, 1, t_gather)559 560 # Build (B, K) labels: 1 at valid aligned positions, -100 (ignore) at561 # sentinel/padding positions so they drop out of the JSD reduction.562 aligned_labels = torch.where(563 aligned_mask,564 torch.ones_like(s_idx),565 torch.full_like(s_idx, -100),566 )567 568 return generalized_jsd_loss(569 student_logits=student_aligned,570 teacher_logits=teacher_aligned,571 labels=aligned_labels, # sentinel-masked aligned error-turn positions572 beta=self.sdpo_jsd_beta,573 temperature=self.sdpo_temperature,574 token_clip=self.sdpo_token_clip,575 reduction="batchmean",576 )577 578 # ----------------------------------------------------------------------579 # Channel 3: trace-replay DPO580 # ----------------------------------------------------------------------581 582 def _compute_trace_replay_loss(583 self,584 model: torch.nn.Module,585 inputs: dict[str, torch.Tensor],586 ) -> torch.Tensor:587 """Standard DPO loss using (chosen, rejected) pairs from teacher disagreement.588 589 DPO loss formula (Rafailov et al. 2023):590 L = -log σ(β · (logπ(chosen) - logπ_ref(chosen)591 - logπ(rejected) + logπ_ref(rejected)))592 593 Where logπ_ref are precomputed by the data collator using the594 reference (init student) policy.595 """596 if (597 self.beta_replay == 0.0598 or "dpo_chosen_input_ids" not in inputs599 or inputs["dpo_chosen_input_ids"].numel() == 0600 ):601 return torch.tensor(0.0, device=_device_of(model), requires_grad=True)602 603 # Forward passes for chosen and rejected, gather logprobs at response tokens604 chosen_logprobs = self._sequence_logprobs(605 model, inputs["dpo_chosen_input_ids"], inputs["dpo_chosen_response_mask"]606 )607 rejected_logprobs = self._sequence_logprobs(608 model, inputs["dpo_rejected_input_ids"], inputs["dpo_rejected_response_mask"]609 )610 611 ref_chosen_logprobs = inputs["dpo_chosen_ref_logprobs"]612 ref_rejected_logprobs = inputs["dpo_rejected_ref_logprobs"]613 614 logits = self.replay_dpo_beta * (615 (chosen_logprobs - ref_chosen_logprobs)616 - (rejected_logprobs - ref_rejected_logprobs)617 )618 return -F.logsigmoid(logits).mean()619 620 @staticmethod621 def _sequence_logprobs(622 model: torch.nn.Module,623 input_ids: torch.Tensor,624 response_mask: torch.Tensor,625 ) -> torch.Tensor:626 """Sum logprob of response tokens given the prompt prefix.627 628 Standard DPO accounting: we only score the response tokens (where629 response_mask == 1), not the prompt tokens.630 """631 outputs = model(input_ids=input_ids)632 # Shift for next-token prediction: logits[t] predicts input_ids[t+1]633 logits = outputs.logits[:, :-1, :]634 targets = input_ids[:, 1:]635 log_probs = F.log_softmax(logits, dim=-1)636 token_logprobs = log_probs.gather(-1, targets.unsqueeze(-1)).squeeze(-1)637 # Mask out prompt + padding; sum response-token logprobs638 masked = token_logprobs * response_mask[:, 1:].float()639 return masked.sum(dim=-1)640 641 642def _device_of(model: torch.nn.Module) -> torch.device:643 """Return the device of any parameter of the model — robust to FSDP/DDP wrappers."""644 return next(model.parameters()).device645 646 647def validate_kl_in_reward_config(648 *,649 kl_estimator: str,650 beta: float,651 scale_rewards: Any,652) -> None:653 """Validate the (kl_estimator, beta, scale_rewards) combo for k1-in-reward.654 655 Extracted so the preconditions are unit-testable without standing up a real656 GRPOTrainer (which needs a model + dataset). Raises ``ValueError`` on any657 invalid combination; returns None when the config is sound.658 659 Preconditions (see ``kl_in_reward.py`` for the algebra):660 * ``kl_estimator`` in {k1, k3}.661 * ``beta != 0`` — TRL only builds the reference model and computes ref662 logprobs when beta>0, and the in-reward penalty needs ref logps. beta663 doubles as the in-reward KL coefficient (the in-loss k3 term is664 suppressed per step).665 * ``scale_rewards`` in {none, false} — the advantage-adjustment identity666 is exact only without per-group std-normalization (the Dr.GRPO /667 Composer regime).668 """669 if kl_estimator not in ("k1", "k3"):670 raise ValueError(f"kl_estimator must be 'k1' or 'k3', got {kl_estimator!r}.")671 if float(beta) == 0.0:672 raise ValueError(673 "kl_in_reward=True requires a non-zero `beta` (the KL coefficient): "674 "TRL only creates the reference model and computes ref logprobs when "675 "beta>0, and k1-in-reward needs those ref logps. Set beta to your KL "676 "coefficient (e.g. make_po_config('dr_grpo', beta=0.04)); the in-loss "677 "k3 term is suppressed automatically so beta acts purely as the "678 "in-reward k1 coefficient."679 )680 if str(scale_rewards).lower() not in ("none", "false"):681 raise ValueError(682 "kl_in_reward=True requires scale_rewards in {none,false} "683 f"(got {scale_rewards!r}). The advantage-adjustment identity "684 "adv -= beta·(KL - group_mean(KL)) is EXACT only without per-group "685 "std-normalization (the Dr.GRPO / Composer regime). With std-norm, "686 "folding KL into the reward also shifts the group std, so the linear "687 "correction no longer matches true in-reward KL. Use "688 "make_po_config('dr_grpo', beta=…) (scale_rewards='none')."689 )690 691 692def make_dr_grpo_config(**overrides: Any):693 """Build a `trl.GRPOConfig` configured to the **Dr. GRPO** recipe.694 695 Per the Composer 2 technical report (arXiv:2603.24477,696 research/10-composer2-techreport-mining.md) the RL base is Dr. GRPO697 (Liu et al., arXiv:2503.20783):698 699 - ``loss_type="dr_grpo"`` — removes GRPO's length-standardization term700 (which injects a length bias). TRL's own help text cites the Dr. GRPO701 paper for this.702 - ``scale_rewards="none"`` — NO std-dev advantage normalization. TRL docs:703 "The Dr. GRPO paper recommends not scaling rewards, as scaling by the704 standard deviation introduces a question-level difficulty bias."705 - ``num_iterations=1`` — single-epoch regime (a prompt is never706 trained on twice), matching the tech report.707 - ``beta`` (KL-to-ref coef) kept. NOTE on the KL estimator (ADR-012708 finding #1, verified against the installed trl==1.5.0 source):709 ``GRPOTrainer._compute_loss`` uses the **k3** estimator710 ``exp(ref_logp - logp) - (ref_logp - logp) - 1``711 (trl/trainer/grpo_trainer.py ~L2513), NOT the k1 estimator712 ``-log r == (ref_logp - logp)``. k3 is Schulman's low-variance,713 always-non-negative KL approximation; k1 is its unbiased but714 higher-variance counterpart. The Dr. GRPO / Composer 2 report discusses715 KL in k1 terms, but the delta is small for r≈1 (k3 = k1 + O((Δlogp)^2))716 and TRL's k3 choice is the production reality. We do NOT monkeypatch TRL717 to force k1; we document the honest delta. See718 ``test_dr_grpo_config_and_alignment.py::test_trl_kl_estimator_is_k3_not_k1``.719 720 Any field can be overridden via kwargs (e.g. ``learning_rate=...``,721 ``output_dir=...``). The three Dr. GRPO-defining knobs are forced unless722 explicitly overridden, and a sanity assertion guards against silent drift.723 """724 from trl import GRPOConfig # local import: only when actually building a config725 726 dr_grpo_defaults: dict[str, Any] = {727 "loss_type": "dr_grpo",728 "scale_rewards": "none",729 "num_iterations": 1,730 }731 merged = {**dr_grpo_defaults, **overrides}732 cfg = GRPOConfig(**merged)733 # Guard: fail loudly if a future TRL renames/repurposes these knobs.734 assert cfg.loss_type == merged["loss_type"], (735 f"GRPOConfig loss_type drifted: requested {merged['loss_type']!r}, "736 f"got {cfg.loss_type!r} — TRL may have renamed/repurposed the knob."737 )738 # Dr. GRPO requires NO std-dev advantage normalization. TRL accepts either739 # the string "none" or the bool False to disable it; normalize before740 # comparing so a future TRL that switches the representation still passes741 # (and a genuinely-wrong value like "batch"/"group"/True fails loudly).742 # (Cross-family review 2026-05-29: the prior literal `("none","False","False")`743 # had a duplicated "False" and did a brittle case-sensitive str compare.)744 assert str(cfg.scale_rewards).lower() in ("none", "false"), (745 f"Dr. GRPO requires scale_rewards disabled (no std-norm); got "746 f"{cfg.scale_rewards!r}. TRL knob may have drifted — re-verify against trl version."747 )748 assert cfg.num_iterations == merged["num_iterations"], "GRPOConfig dropped num_iterations"749 return cfg750 751 752# ---------------------------------------------------------------------------753# Policy-optimization objective MENU (ADR-014)754# ---------------------------------------------------------------------------755#756# The base RL objective used to be hardcoded to Dr.GRPO (make_dr_grpo_config).757# make_po_config gives RL a real menu: GRPO-family objectives selectable by name.758# Verified against the installed trl==1.5.0 (introspected 2026-05-30): its759# GRPOTrainer already implements these as `loss_type` branches + knobs, so EVERY760# preset below is pure config — no custom _compute_loss override needed.761#762# Knob-space each preset sets (all real GRPOConfig fields in trl 1.5.0):763# loss_type ∈ {grpo, dr_grpo, bnpo, dapo, cispo} (gspo = grpo loss +764# importance_sampling_level="sequence"; trl has no literal "gspo")765# scale_rewards ∈ {"group"(std-norm), "batch", "none"(no std-norm, Dr.GRPO)}766# epsilon / epsilon_high — symmetric vs decoupled "clip-higher" (DAPO)767# importance_sampling_level ∈ {"token", "sequence"(GSPO)}768# beta — KL-to-ref coef (0.0 = reference-free)769# mask_truncated_completions — DAPO overlong masking770# num_iterations — on-policy reuse (1 = strict on-policy)771 772#: Selectable base policy-optimization objectives (named presets over trl knobs).773PO_OBJECTIVES: dict[str, dict[str, Any]] = {774 # Vanilla GRPO (DeepSeekMath, arXiv 2402.03300): group-relative advantage775 # WITH std normalization + per-sequence length normalization, KL on.776 "grpo": {777 "loss_type": "grpo",778 "scale_rewards": "group",779 "importance_sampling_level": "token",780 "num_iterations": 1,781 },782 # Dr.GRPO (arXiv 2503.20783): remove length-std normalization bias (no783 # advantage /std, length-independent aggregation). Framework's historical784 # default (== make_dr_grpo_config). Composer 2.5's base objective.785 "dr_grpo": {786 "loss_type": "dr_grpo",787 "scale_rewards": "none",788 "importance_sampling_level": "token",789 "num_iterations": 1,790 },791 # BNPO: batch-normalized variant (trl loss_type), std over the batch.792 "bnpo": {793 "loss_type": "bnpo",794 "scale_rewards": "batch",795 "importance_sampling_level": "token",796 "num_iterations": 1,797 },798 # DAPO (arXiv 2503.14476): decoupled "clip-higher" (epsilon_high > epsilon)799 # + token-level loss + overlong masking + KL removed. High-value, low-cost800 # anti-entropy-collapse objective. epsilon_high=0.28 per the paper.801 "dapo": {802 "loss_type": "dapo",803 "scale_rewards": "none",804 "epsilon": 0.2,805 "epsilon_high": 0.28,806 "mask_truncated_completions": True,807 "beta": 0.0,808 "importance_sampling_level": "token",809 "num_iterations": 1,810 },811 # GSPO (Qwen, arXiv 2507.18071): SEQUENCE-level importance ratio (one length-812 # normalized ratio per response) — stabilizes long-CoT and especially MoE RL.813 # trl expresses this as the grpo loss + importance_sampling_level="sequence".814 "gspo": {815 "loss_type": "grpo",816 "scale_rewards": "group",817 "importance_sampling_level": "sequence",818 "num_iterations": 1,819 },820 # CISPO (MiniMax-M1, arXiv 2506.13585): clip the IS weight and detach it as a821 # constant coefficient on log π — every token keeps a gradient (fixes the822 # "rare reasoning tokens get zeroed by the clip" pathology). eps_max≈5 (ScaleRL).823 "cispo": {824 "loss_type": "cispo",825 "scale_rewards": "none",826 "epsilon_high": 5.0,827 "importance_sampling_level": "token",828 "num_iterations": 1,829 },830}831 832 833def make_po_config(objective: str = "dr_grpo", **overrides: Any):834 """Build a `trl.GRPOConfig` for a NAMED policy-optimization objective.835 836 The menu that gives RL real options beyond the single hardcoded Dr.GRPO837 recipe. ``objective`` selects a preset from ``PO_OBJECTIVES`` (grpo /838 dr_grpo / bnpo / dapo / gspo / cispo); ``**overrides`` set or override any839 GRPOConfig field on top (e.g. ``output_dir=...``, ``beta=...``,840 ``learning_rate=...``).841 842 All presets are PURE CONFIG over trl 1.5.0's GRPOTrainer (verified by843 introspecting the installed package 2026-05-30): the trainer already844 implements each ``loss_type`` branch and the ``importance_sampling_level`` /845 ``epsilon_high`` knobs, so no custom ``_compute_loss`` is needed. See ADR-014.846 847 Raises:848 ValueError: unknown objective (lists the valid menu).849 AssertionError: a requested knob silently failed to apply (drift guard).850 """851 from trl import GRPOConfig # local import: only when actually building a config852 853 key = (objective or "dr_grpo").lower()854 if key not in PO_OBJECTIVES:855 raise ValueError(856 f"Unknown PO objective {objective!r}. Choose from: "857 f"{sorted(PO_OBJECTIVES)}. (Each is a named preset over trl 1.5.0's "858 f"GRPOConfig knobs — see PO_OBJECTIVES / ADR-014.)"859 )860 861 preset = dict(PO_OBJECTIVES[key])862 merged = {**preset, **overrides}863 cfg = GRPOConfig(**merged)864 865 # Drift guards: fail loudly if a future trl renamed/repurposed a knob we set,866 # so a preset can never silently degrade to a different objective.867 if "loss_type" in merged:868 assert str(cfg.loss_type) == str(merged["loss_type"]), (869 f"GRPOConfig.loss_type drifted: requested {merged['loss_type']!r}, "870 f"got {cfg.loss_type!r} — trl may have renamed the knob."871 )872 if "importance_sampling_level" in merged and hasattr(cfg, "importance_sampling_level"):873 assert str(cfg.importance_sampling_level) == str(874 merged["importance_sampling_level"]875 ), (876 f"importance_sampling_level drifted for objective {key!r}: requested "877 f"{merged['importance_sampling_level']!r}, got {cfg.importance_sampling_level!r}."878 )879 if key == "gspo":880 assert str(getattr(cfg, "importance_sampling_level", "token")) == "sequence", (881 "GSPO requires importance_sampling_level='sequence'; it was overridden "882 "to token, which silently degrades GSPO to GRPO. Drop that override."883 )884 if merged.get("epsilon_high") is not None:885 assert abs(886 float(getattr(cfg, "epsilon_high", merged["epsilon_high"]))887 - float(merged["epsilon_high"])888 ) < 1e-9, f"epsilon_high (decoupled clip) drifted for {key!r}."889 return cfg890 891 892__all__ = [893 "ComposerReplicationTrainer",894 "make_dr_grpo_config",895 "make_po_config",896 "PO_OBJECTIVES",897 "validate_kl_in_reward_config",898]899 