Download loopq_quantization/scripts/loopq/objective.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 9.48 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/objective.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/objective.py
-
curl -L -o objective.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/objective.py
9.48 kB
| """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") | |
| 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") | |