"""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")