Codeseys/composer-replication-framework
0
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 