Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
composer_trainer.py899 linesDownload Raw Back to trainer
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