File size: 9,479 Bytes
9118991
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
"""Trajectory-aware LoopQ calibration objective (Eq. 8 and Appendix B.4)."""

from __future__ import annotations

from dataclasses import dataclass

import torch
import torch.nn.functional as F


PAPER_TRAJECTORY_LAMBDA = 0.1
PAPER_KL_TEMPERATURE = 1.0
PAPER_TEACHER_TOP_K = 1000
PAPER_EXAMPLE_MU_UPDATE_INTERVAL = 100
DEFAULT_MU_EPSILON = 1e-8
KL_MODES = ("conditional_topk", "topk_with_tail", "full")


@dataclass(frozen=True)
class TrajectoryLoss:
    total: torch.Tensor
    kl: torch.Tensor
    same_loop: torch.Tensor
    final_teacher: torch.Tensor
    transition: torch.Tensor
    trajectory: torch.Tensor
    mu: torch.Tensor
    effective_top_k: int

    def trace(self) -> dict[str, float | int | list[float]]:
        return {
            "total": float(self.total.detach()),
            "kl": float(self.kl.detach()),
            "same_loop": float(self.same_loop.detach()),
            "final_teacher": float(self.final_teacher.detach()),
            "transition": float(self.transition.detach()),
            "trajectory": float(self.trajectory.detach()),
            "mu": [float(value) for value in self.mu.detach().cpu()],
            "effective_top_k": self.effective_top_k,
        }


class AdaptiveMuCache:
    """Periodically refreshed, detached Appendix B.4 trust weights."""

    def __init__(self, update_interval: int = PAPER_EXAMPLE_MU_UPDATE_INTERVAL) -> None:
        if update_interval <= 0:
            raise ValueError("mu update_interval must be positive")
        self.update_interval = int(update_interval)
        self.last_update_step: int | None = None
        self.value: torch.Tensor | None = None

    def get(
        self,
        step: int,
        teacher_hidden: torch.Tensor,
        student_hidden: torch.Tensor,
        *,
        epsilon: float = DEFAULT_MU_EPSILON,
    ) -> torch.Tensor:
        if step < 0:
            raise ValueError("calibration step must be non-negative")
        should_update = self.value is None or (step % self.update_interval == 0 and step != self.last_update_step)
        if should_update:
            self.value = adaptive_mu(
                teacher_hidden.detach(), student_hidden.detach(), epsilon=epsilon
            ).cpu()
            self.last_update_step = step
        return self.value.to(device=student_hidden.device, dtype=torch.float32)

    def state_dict(self) -> dict[str, object]:
        return {
            "update_interval": self.update_interval,
            "last_update_step": self.last_update_step,
            "value": None if self.value is None else self.value.clone(),
        }

    def load_state_dict(self, state: dict[str, object]) -> None:
        if int(state["update_interval"]) != self.update_interval:
            raise ValueError("mu update interval does not match checkpoint")
        self.last_update_step = state["last_update_step"]
        value = state["value"]
        self.value = None if value is None else torch.as_tensor(value).clone()


def adaptive_mu(
    teacher_hidden: torch.Tensor,
    student_hidden: torch.Tensor,
    *,
    epsilon: float = DEFAULT_MU_EPSILON,
) -> torch.Tensor:
    """Compute Appendix B.4 mu_t from complete recurrent trajectories."""

    _validate_hidden_trajectories(teacher_hidden, student_hidden)
    if epsilon <= 0:
        raise ValueError("mu epsilon must be positive")
    teacher = teacher_hidden.detach().to(torch.float64)
    student = student_hidden.detach().to(torch.float64)
    final_teacher = teacher[-1]
    final_distances = (teacher - final_teacher).square().flatten(1).sum(dim=1)
    mismatches = (teacher - student).square().flatten(1).sum(dim=1)
    suffix_mismatch = torch.flip(torch.cumsum(torch.flip(mismatches, dims=(0,)), dim=0), dims=(0,))
    return (final_distances / (final_distances + suffix_mismatch + epsilon)).to(torch.float32)


def topk_teacher_kl(
    teacher_logits: torch.Tensor,
    student_logits: torch.Tensor,
    *,
    top_k: int = PAPER_TEACHER_TOP_K,
    temperature: float = PAPER_KL_TEMPERATURE,
    mode: str = "conditional_topk",
) -> tuple[torch.Tensor, int]:
    """Explicit alternatives for the paper's unspecified top-k tail treatment.

    conditional_topk preserves the original local objective. topk_with_tail
    aggregates all excluded vocabulary items into one category, retaining its
    probability mass. full is a diagnostic, not the paper's top-1000 setting.
    All modes sum tokens/classes and average only the leading batch dimension.
    """
    if teacher_logits.shape != student_logits.shape or teacher_logits.ndim < 2:
        raise ValueError("teacher and student logits must have the same rank>=2 shape")
    if top_k <= 0 or temperature <= 0:
        raise ValueError("top_k and temperature must be positive")
    if mode not in KL_MODES:
        raise ValueError(f"KL mode must be one of {KL_MODES}")
    if mode == "full":
        teacher_logp = (teacher_logits.detach().float() / temperature).log_softmax(-1)
        student_logp = (student_logits.float() / temperature).log_softmax(-1)
        return (F.kl_div(student_logp, teacher_logp, log_target=True, reduction="batchmean")
                * temperature * temperature, teacher_logits.shape[-1])
    effective_top_k = min(top_k, teacher_logits.shape[-1])
    indices = teacher_logits.detach().topk(effective_top_k, dim=-1).indices
    teacher_selected = teacher_logits.detach().gather(-1, indices).float() / temperature
    student_selected = student_logits.gather(-1, indices).float() / temperature
    if mode == "topk_with_tail" and effective_top_k < teacher_logits.shape[-1]:
        # logsumexp of excluded logits avoids unstable 1 - sum(top-k probs).
        mask = torch.zeros_like(teacher_logits, dtype=torch.bool).scatter_(-1, indices, True)
        teacher_tail = (teacher_logits.detach().float() / temperature).masked_fill(mask, -torch.inf).logsumexp(-1, keepdim=True)
        student_tail = (student_logits.float() / temperature).masked_fill(mask, -torch.inf).logsumexp(-1, keepdim=True)
        teacher_selected = torch.cat((teacher_selected, teacher_tail), dim=-1)
        student_selected = torch.cat((student_selected, student_tail), dim=-1)
    teacher_probability = teacher_selected.softmax(dim=-1)
    student_log_probability = student_selected.log_softmax(dim=-1)
    # Standard distillation scaling is neutral at the paper's temperature 1.
    kl = F.kl_div(student_log_probability, teacher_probability, reduction="batchmean")
    return kl * (temperature * temperature), effective_top_k


def trajectory_aware_loss(
    *,
    teacher_logits: torch.Tensor,
    student_logits: torch.Tensor,
    teacher_hidden: torch.Tensor,
    student_hidden: torch.Tensor,
    adapted_transitions: torch.Tensor,
    teacher_next_inputs: torch.Tensor,
    mu: torch.Tensor,
    include_transition: bool = True,
    trajectory_lambda: float = PAPER_TRAJECTORY_LAMBDA,
    top_k: int = PAPER_TEACHER_TOP_K,
    temperature: float = PAPER_KL_TEMPERATURE,
    kl_mode: str = "conditional_topk",
) -> TrajectoryLoss:
    """Evaluate LoopQ's practical Appendix B.4 form of Equation (8)."""

    _validate_hidden_trajectories(teacher_hidden, student_hidden)
    loops = teacher_hidden.shape[0]
    expected_transition_shape = (loops - 1,) + tuple(teacher_hidden.shape[1:])
    if tuple(adapted_transitions.shape) != expected_transition_shape:
        raise ValueError(f"adapted_transitions must have shape {expected_transition_shape}")
    if tuple(teacher_next_inputs.shape) != expected_transition_shape:
        raise ValueError(f"teacher_next_inputs must have shape {expected_transition_shape}")
    if tuple(mu.shape) != (loops,) or (mu < 0).any() or (mu > 1).any():
        raise ValueError(f"mu must have shape ({loops},) with values in [0, 1]")
    if trajectory_lambda < 0:
        raise ValueError("trajectory_lambda must be non-negative")

    kl, effective_top_k = topk_teacher_kl(
        teacher_logits, student_logits, top_k=top_k, temperature=temperature, mode=kl_mode
    )
    student_hidden = student_hidden.float()
    teacher = teacher_hidden.detach().to(student_hidden)
    final_teacher = teacher[-1]
    same_per_loop = (student_hidden - teacher).square().flatten(1).sum(dim=1)
    final_per_loop = (student_hidden - final_teacher).square().flatten(1).sum(dim=1)
    mu_typed = mu.detach().to(student_hidden)
    same_weighted = ((1 - mu_typed) * same_per_loop).sum()
    final_weighted = (mu_typed * final_per_loop).sum()
    transition = (
        adapted_transitions.float() - teacher_next_inputs.detach().to(device=adapted_transitions.device, dtype=torch.float32)
    ).square().flatten(1).sum(dim=1).sum()
    if not include_transition:
        transition = transition.new_zeros(())
    trajectory = same_weighted + final_weighted + transition
    total = kl + trajectory_lambda * trajectory
    return TrajectoryLoss(
        total=total,
        kl=kl,
        same_loop=same_weighted,
        final_teacher=final_weighted,
        transition=transition,
        trajectory=trajectory,
        mu=mu_typed,
        effective_top_k=effective_top_k,
    )


def _validate_hidden_trajectories(
    teacher_hidden: torch.Tensor, student_hidden: torch.Tensor
) -> None:
    if teacher_hidden.shape != student_hidden.shape or teacher_hidden.ndim < 2:
        raise ValueError("teacher and student hidden trajectories must have the same rank>=2 shape")
    if teacher_hidden.shape[0] < 2:
        raise ValueError("trajectory must contain at least two loops")