Codeseys/composer-replication-framework
0
1"""TAID loss — Temporally Adaptive Interpolated Distillation.2 3Paper: "TAID: Temporally Adaptive Interpolated Distillation for Efficient4 Knowledge Transfer in Language Models"5 Sakana AI, arXiv:2501.169376License: Apache-2.0 (https://github.com/SakanaAI/TAID)7 8This module is a faithful port of the reference implementation at9``SakanaAI/TAID/src/distil_losses/taid.py``. **The previous in-tree10implementation was algorithmically different from the paper** (it mixed in11probability space against a frozen step-0 student snapshot and wrapped a12symmetric JSD criterion). This rewrite replaces it with the upstream13algorithm:14 15 p_t = softmax( (1 - t) · stop_grad(student_logits) + t · teacher_logits )16 loss = - mean_token Σ_v p_t(v) · log_softmax(student_logits)(v)17 18That is:19 1. Mix in **logit space**, not probability space.20 2. Anchor against the **current student detached** (re-evaluated each21 step), not a frozen step-0 snapshot.22 3. Distillation criterion is **forward KL** (Hinton-style soft target),23 not symmetric JSD.24 25Schedule26--------27The original implementation embedded an adaptive momentum-based schedule28inside the loss object; this is now factored out into the optional29:class:`TAIDScheduler` so the loss function itself is pure (single ``t``30in [0, 1]). Callers either:31 32- Pass a fixed ``t`` for ablations / fixed schedules.33- Drive ``t`` via :class:`TAIDScheduler` (paper-default adaptive scheme).34- Drive ``t`` via any custom schedule of their choosing.35 36Backward-incompatible change37----------------------------38The previous public signature was:39 40 taid_loss(student_logits, teacher_logits, student_init_logits, *,41 schedule_step, total_steps, schedule, alpha_min, alpha_max,42 jsd_beta, temperature, reduction)43 44The new signature is:45 46 taid_loss(student_logits, teacher_logits, mask=None, *, t)47 48Removed kwargs (``student_init_logits``, ``schedule_step``, ``total_steps``,49``schedule``, ``alpha_min``, ``alpha_max``, ``jsd_beta``, ``temperature``,50``reduction``) have no upstream analogue. Pass ``t`` directly; if you need51a schedule, use :class:`TAIDScheduler` or compute ``t`` yourself.52 53Reference: arXiv:2501.16937; ``SakanaAI/TAID`` commit history.54"""55from __future__ import annotations56 57import torch58import torch.nn.functional as F59 60 61def taid_loss(62 student_logits: torch.Tensor,63 teacher_logits: torch.Tensor,64 mask: torch.Tensor | None = None,65 *,66 t: float | torch.Tensor,67) -> torch.Tensor:68 """TAID forward-KL loss against a logit-space-interpolated target.69 70 Faithful port of ``SakanaAI/TAID/src/distil_losses/taid.py:compute_loss``71 composed with ``fkl.forward_kl``.72 73 Pseudocode::74 75 p_t = softmax( (1 - t) · student_logits.detach() + t · teacher_logits )76 log_q = log_softmax( student_logits )77 per_token = - Σ_v p_t(v) · log_q(v) # forward KL token-wise78 loss = sum(per_token · mask) / sum(mask)79 80 Args:81 student_logits: ``(B, T, V)`` current student logits, with grad.82 teacher_logits: ``(B, T, V)`` teacher logits (no grad expected;83 detached internally only insofar as the interpolation uses the84 student detach — teacher gradient is left untouched, matching85 upstream).86 mask: ``(B, T)`` token mask (1 = include, 0 = ignore). Required by87 upstream; defaults to all-ones if omitted for convenience.88 t: interpolation coefficient in ``[0, 1]``. Scalar Python float or89 0-d torch.Tensor. ``t=0`` makes the target match the (detached)90 student — a regularizer with zero gradient signal. ``t=1`` makes91 the target the teacher — pure forward-KL distillation.92 93 Returns:94 Scalar loss (token-mean, in float32 dtype matching upstream).95 96 Raises:97 ValueError: shape mismatch between student/teacher, or invalid mask98 shape.99 100 Reference: arXiv:2501.16937 §3.1 + Eq. (4); upstream commit at101 ``SakanaAI/TAID@main:src/distil_losses/taid.py``.102 """103 if student_logits.shape != teacher_logits.shape:104 raise ValueError(105 f"student/teacher logits shape mismatch: "106 f"{tuple(student_logits.shape)} vs {tuple(teacher_logits.shape)}"107 )108 if mask is None:109 mask = student_logits.new_ones(student_logits.shape[:-1])110 elif mask.shape != student_logits.shape[:-1]:111 raise ValueError(112 f"mask shape {tuple(mask.shape)} does not match logits prefix "113 f"{tuple(student_logits.shape[:-1])}"114 )115 116 # 1. Logit-space mix with student detached (anchor = current student, no grad).117 blended_logits = (1 - t) * student_logits.detach() + t * teacher_logits118 119 # 2. Target distribution in float32 for numerical stability (upstream choice).120 p_t = F.softmax(blended_logits, dim=-1, dtype=torch.float32)121 122 # 3. Forward KL: the gradient flows ONLY through student log-softmax.123 student_logprobs = F.log_softmax(student_logits, dim=-1, dtype=torch.float32)124 125 # 4. Mask out -inf positions in the student logits (upstream guard).126 inf_mask = torch.isinf(student_logits)127 prod = torch.masked_fill(p_t * student_logprobs, inf_mask, 0.0)128 129 # 5. Per-token cross-entropy = -sum_v p_t(v) * log_q(v); reduce over vocab.130 per_token = -prod.sum(dim=-1).reshape(-1)131 flat_mask = mask.reshape(-1).to(per_token.dtype)132 denom = flat_mask.sum().clamp_min(1.0)133 loss = (per_token * flat_mask).sum() / denom134 return loss135 136 137class TAIDScheduler:138 """Adaptive momentum-based schedule for TAID's interpolation coefficient ``t``.139 140 Stateful, mirrors ``SakanaAI/TAID/src/distil_losses/taid.py:TAID.update_t``.141 142 Usage::143 144 sched = TAIDScheduler(num_train_steps=10_000)145 for step in range(num_train_steps):146 t = sched.t # current t (float)147 loss = taid_loss(s_logits, t_logits, mask, t=t)148 loss.backward(); optimizer.step()149 sched.update_t(loss.detach(), global_step=step)150 151 The schedule is monotone non-decreasing: at each step, the floor is the152 linear schedule ``t_target = t_start + progress · (t_end - t_start)``,153 and an adaptive bump ``alpha · σ(momentum) · (1 - t)`` is added on top154 where ``momentum`` tracks the relative loss change with EMA decay155 ``beta``. ``disable_adaptive=True`` collapses to the deterministic linear156 schedule.157 158 Args:159 num_train_steps: total planned training steps; required so the linear160 floor ``t_target`` is well-defined.161 t_start: initial ``t`` (paper default 0.4 — the student is already162 close to the teacher in this regime, so ``t=0`` would waste the163 warmup phase).164 t_end: terminal ``t`` (paper default 1.0).165 alpha: adaptive bump magnitude (paper default 5e-4).166 beta: EMA decay for the relative-loss-change momentum (paper default167 0.99).168 disable_adaptive: if True, fall back to deterministic linear schedule169 ``t_target = t_start + progress · (t_end - t_start)``.170 device: device to allocate state buffers on; default cpu.171 """172 173 def __init__(174 self,175 num_train_steps: int,176 *,177 t_start: float = 0.4,178 t_end: float = 1.0,179 alpha: float = 5e-4,180 beta: float = 0.99,181 disable_adaptive: bool = False,182 device: torch.device | str = "cpu",183 ) -> None:184 if not (0.0 <= t_start < 1.0):185 raise ValueError(f"t_start must be in [0, 1), got {t_start}")186 if not (0.0 < t_end <= 1.0):187 raise ValueError(f"t_end must be in (0, 1], got {t_end}")188 if not (0.0 <= alpha <= 1.0):189 raise ValueError(f"alpha must be in [0, 1], got {alpha}")190 if num_train_steps <= 0:191 raise ValueError(f"num_train_steps must be > 0, got {num_train_steps}")192 193 self.t_start = t_start194 self.t_end = t_end195 self.alpha = alpha196 self.beta = beta197 self.disable_adaptive = disable_adaptive198 self.num_train_steps = num_train_steps199 200 self._t = torch.tensor(t_start, device=device, dtype=torch.float32)201 self._prev_loss = torch.tensor(202 float("inf"), device=device, dtype=torch.float32203 )204 self._momentum = torch.zeros([], device=device, dtype=torch.float32)205 206 @property207 def t(self) -> float:208 """Current interpolation coefficient as a Python float."""209 return float(self._t)210 211 def update_t(212 self,213 loss: torch.Tensor,214 global_step: int,215 ) -> torch.Tensor | None:216 """Update internal ``t`` given the current step's distillation loss.217 218 Mirrors upstream verbatim. First call with finite loss only seeds219 ``prev_loss`` and returns None. Subsequent calls update momentum +220 ``t`` and return the (positive) ``delta_t`` that was added on top of221 the linear floor (None for the first call).222 223 Args:224 loss: scalar loss tensor (caller should pass ``loss.detach()``).225 global_step: current global step (0-indexed).226 227 Returns:228 The adaptive ``delta_t`` that was applied, or None if this was229 the seeding call.230 """231 if torch.isinf(self._prev_loss):232 self._prev_loss = loss.detach().to(self._prev_loss)233 return None234 235 relative_change = (self._prev_loss - loss) / (self._prev_loss + 1e-15)236 self._momentum = (237 self.beta * self._momentum + (1 - self.beta) * relative_change238 )239 240 adaptive_delta = torch.sigmoid(self._momentum)241 progress = global_step / self.num_train_steps242 t_target = self.t_start + (self.t_end - self.t_start) * progress243 delta_t = self.alpha * adaptive_delta * (1 - self._t)244 245 if self.disable_adaptive:246 new_t = t_target247 else:248 new_t = min(self.t_end, max(t_target, float(self._t + delta_t)))249 250 if not isinstance(new_t, torch.Tensor):251 new_t = torch.tensor(new_t, device=self._t.device, dtype=self._t.dtype)252 self._t = new_t253 self._prev_loss = loss.detach().to(self._prev_loss)254 return delta_t255 256 257__all__ = ["taid_loss", "TAIDScheduler"]258 