Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
taid.py258 linesDownload Raw Back to distillation
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