JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
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")
@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")