Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
1"""opsd_loss.py — Self-distillation loss, lifted from siyan-zhao/OPSD.2 3Original source: github.com/siyan-zhao/OPSD::OPSDTrainer.generalized_jsd_loss (MIT).4Verified self-contained via DeepWiki audit on 2026-05-25.5Re-aligned byte-for-byte against upstream `opsd_trainer.py` lines 381-479 on62026-05-26 after Wave 15 math review found three numerical divergences (mixture7weighting, β coefficient placement, reduction divisor) and one docstring mislabel.8 9Mathematical reference:10- OPSD paper: Zhao et al., "Self-Distilled Reasoner: On-Policy Self-Distillation11  for LLMs", arXiv:2601.18734.12- SDPO paper: Hübotter et al., "Reinforcement Learning via Self-Distillation",13  arXiv:2601.20802. PROVENANCE (corrected per deepread finding V1): Cursor's14  blog cites SDPO/OPSD only as *background* ("For more background on this15  approach see…"), NOT as its mechanism. Published SDPO distills over the FULL16  rollout with feedback in the prefix and an EMA-regularized teacher; this17  repo's channel is a turn-localized hint-splice with a live (stop-grad,18  non-EMA) teacher — a third, blog-inspired design, neither verbatim SDPO nor19  confirmed-Composer. The kernel below matches OPSD's generalized JSD math.20 21The loss computes JSD/KL divergence between a teacher distribution (model22conditioned on privileged information / a hint) and a student distribution23(model on the original context). Both come from the SAME model — the teacher24is just "the model with hint inserted into context."25 26Composer 2.5's blog describes inserting a "hint" at the error-turn site and27distilling the student toward the hint-conditioned distribution "for that turn28only". The data collator constructs ctx_teacher = ctx_student +29hint_at_error_turn for us.30"""31 32from __future__ import annotations33 34import torch35import torch.nn.functional as F36 37 38def generalized_jsd_loss(39    student_logits: torch.Tensor,40    teacher_logits: torch.Tensor,41    labels: torch.Tensor | None = None,42    beta: float = 0.5,43    temperature: float = 1.0,44    reduction: str = "batchmean",45    logits_are_probs: bool = False,46    top_k: int | None = None,47    token_clip: float | None = None,48) -> torch.Tensor:49    """Generalized Jensen-Shannon Divergence loss between student and teacher.50 51    Byte-for-byte replication of `OPSDTrainer.generalized_jsd_loss`52    (siyan-zhao/OPSD, opsd_trainer.py lines 381-479). See53    https://huggingface.co/papers/2306.13649 Eq. (1) for the definition.54 55    Args:56        student_logits: (B, T, V) — student model logits at each token position.57        teacher_logits: (B, T, V) — teacher (= same model with hint context) logits.58        labels: (B, T) — token-level mask. Positions with label == -100 are ignored59            (standard HF padding/ignored convention). For Composer-style hint-distill,60            mask should be 1 at error-turn tokens AFTER the hint, 0 elsewhere.61        beta: in [0, 1]. NOTE on direction (per `F.kl_div` semantics, where62            `F.kl_div(log_q, log_p, log_target=True)` computes KL(p || q)):63              β = 0  → kl_div(student_log_probs, teacher_log_probs)64                    = KL(teacher || student)  (reverse KL — mode-covering for student)65              β = 1  → kl_div(teacher_log_probs, student_log_probs)66                    = KL(student || teacher)  (forward KL — mode-seeking for student)67              β = 0.5 → symmetric JSD with M = 0.5*(P+Q)68            General β ∈ (0,1): mixture M = (1-β)·P_student + β·P_teacher and69            jsd = β·KL(teacher||M) + (1-β)·KL(student||M).70        temperature: softens distributions; T > 1 encourages distribution-matching71            on broader tail probabilities. SDPO paper uses 1.0.72        reduction: "batchmean" | "sum" | "mean" | "none". "batchmean" matches73            upstream OPSD: divides by `mask.sum()` when labels are given, else74            by the leading dim of jsd (= batch size). This differs from PyTorch's75            `KLDivLoss(reduction='batchmean')` (which divides by batch). We match76            upstream because gradient scale stability matters more than the name.77        logits_are_probs: if True, inputs are already probabilities (skip softmax).78        top_k: restrict KL to top-k tokens of the teacher distribution.79            Saves compute on large vocabularies (Qwen3 vocab = 152K).80        token_clip: clip per-token JSD to this max. Stabilizes training.81            SDPO paper does NOT clip; OPSD code defaults to None (no clip).82 83    Returns:84        Scalar loss tensor (or unreduced (B, T, V) tensor for reduction="none").85    """86    # Path A: probabilities-in. Take log directly with a clamp for stability.87    if logits_are_probs:88        student_log_probs = torch.log(student_logits.clamp_min(1e-8))89        teacher_log_probs = torch.log(teacher_logits.clamp_min(1e-8))90    else:91        # Apply temperature scaling to logits before computing probabilities.92        student_logits = student_logits / temperature93        teacher_logits = teacher_logits / temperature94 95        if top_k is not None and top_k > 0:96            # Restrict to top-k tokens of the teacher distribution and renormalize.97            _, top_k_indices = torch.topk(teacher_logits, k=top_k, dim=-1)98            student_logits = torch.gather(student_logits, dim=-1, index=top_k_indices)99            teacher_logits = torch.gather(teacher_logits, dim=-1, index=top_k_indices)100 101        student_log_probs = F.log_softmax(student_logits, dim=-1)102        teacher_log_probs = F.log_softmax(teacher_logits, dim=-1)103 104    if beta == 0:105        # F.kl_div(input=log_q, target=log_p, log_target=True) computes KL(p || q):106        #   sum_x p(x) * (log p(x) - log q(x))107        # With input=student_log_probs, target=teacher_log_probs → KL(teacher || student).108        jsd = F.kl_div(student_log_probs, teacher_log_probs, reduction="none", log_target=True)109    elif beta == 1:110        jsd = F.kl_div(teacher_log_probs, student_log_probs, reduction="none", log_target=True)111    else:112        # Compute the log of the β-weighted mixture distribution:113        #   M = (1-β)·P_student + β·P_teacher114        #   log M = logsumexp([log P_student + log(1-β), log P_teacher + log(β)])115        beta = torch.tensor(beta, dtype=student_log_probs.dtype, device=student_log_probs.device)116        mixture_log_probs = torch.logsumexp(117            torch.stack([student_log_probs + torch.log1p(-beta), teacher_log_probs + torch.log(beta)]),118            dim=0,119        )120 121        # Compute KL divergences using F.kl_div.122        # PyTorch differs from the standard mathematical definition, so the order of123        # the probability distributions is swapped compared to that defined in the paper.124        kl_teacher = F.kl_div(mixture_log_probs, teacher_log_probs, reduction="none", log_target=True)125        kl_student = F.kl_div(mixture_log_probs, student_log_probs, reduction="none", log_target=True)126 127        # Generalized JSD: β weights the teacher-leg KL (matches upstream).128        jsd = beta * kl_teacher + (1 - beta) * kl_student129 130    # Per-token clipping: cap each token's divergence value.131    if token_clip is not None:132        jsd = jsd.clamp(max=token_clip)133 134    # Masking. labels has shape (B, T); jsd has shape (B, T, V) (or top_k for V).135    # `jsd[mask]` indexes the first two dims, yielding shape (n_valid, V).136    mask = None137    if labels is not None:138        mask = labels != -100139        jsd = jsd[mask]140 141    # Apply reduction (matches upstream byte-for-byte for batchmean/sum/mean).142    if reduction == "batchmean":143        if labels is not None:144            assert mask is not None145            return jsd.sum() / mask.sum()146        return jsd.sum() / jsd.size(0)147    elif reduction == "sum":148        return jsd.sum()149    elif reduction == "mean":150        return jsd.mean()151    elif reduction == "none":152        return jsd153    else:154        # Upstream falls through to `return jsd` for unknown reductions; we raise155        # to surface caller bugs instead of silently returning an unreduced tensor.156        raise ValueError(f"Unknown reduction: {reduction}")157 158 159__all__ = ["generalized_jsd_loss"]160