Download loopq_quantization/scripts/loopq/cta.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 5.89 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/cta.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/cta.py
-
curl -L -o cta.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/cta.py
5.89 kB
| """Cross-loop Transition Adapter (CTA) from LoopQ Equation (7).""" | |
| from __future__ import annotations | |
| import math | |
| from collections.abc import Mapping | |
| from typing import Any | |
| import torch | |
| from torch import nn | |
| DEFAULT_CTA_RANK = 8 | |
| DEFAULT_RMSNORM_EPSILON = 1e-6 | |
| OURO_TRANSITION_COUNT = 3 | |
| HUGINN_TRANSITION_COUNT = 31 | |
| class CrossLoopTransitionAdapter(nn.Module): | |
| """Apply LoopQ CTA only at a true cross-loop transition. | |
| ``U`` and ``V`` are shared across all transitions. The affine vectors | |
| ``a_t``, ``b_t`` and low-rank gate ``eta_t`` are transition-dependent. | |
| Identity initialization is exact because ``a=1``, ``b=0`` and ``eta=0``. | |
| """ | |
| FORMAT_VERSION = 1 | |
| def __init__( | |
| self, | |
| hidden_size: int, | |
| transition_count: int, | |
| *, | |
| rank: int = DEFAULT_CTA_RANK, | |
| rmsnorm_epsilon: float = DEFAULT_RMSNORM_EPSILON, | |
| dtype: torch.dtype = torch.float32, | |
| ) -> None: | |
| super().__init__() | |
| if hidden_size <= 0: | |
| raise ValueError("hidden_size must be positive") | |
| if transition_count <= 0: | |
| raise ValueError("transition_count must be positive") | |
| if rank <= 0 or rank > hidden_size: | |
| raise ValueError("rank must be positive and no larger than hidden_size") | |
| if not math.isfinite(rmsnorm_epsilon) or rmsnorm_epsilon <= 0: | |
| raise ValueError("rmsnorm_epsilon must be finite and positive") | |
| self.enabled = True | |
| self.hidden_size = int(hidden_size) | |
| self.transition_count = int(transition_count) | |
| self.rank = int(rank) | |
| self.rmsnorm_epsilon = float(rmsnorm_epsilon) | |
| # Deterministic full-rank column initialization. eta=0 keeps the CTA | |
| # exactly identity while leaving a nonzero gradient path to its gates. | |
| basis = torch.eye(hidden_size, dtype=dtype)[:, :rank] | |
| self.U = nn.Parameter(basis.clone()) | |
| self.V = nn.Parameter(basis.clone()) | |
| self.a = nn.Parameter(torch.ones(transition_count, hidden_size, dtype=dtype)) | |
| self.b = nn.Parameter(torch.zeros(transition_count, hidden_size, dtype=dtype)) | |
| self.eta = nn.Parameter(torch.zeros(transition_count, rank, dtype=dtype)) | |
| def for_ouro( | |
| cls, | |
| hidden_size: int, | |
| *, | |
| rank: int = DEFAULT_CTA_RANK, | |
| rmsnorm_epsilon: float = DEFAULT_RMSNORM_EPSILON, | |
| ) -> "CrossLoopTransitionAdapter": | |
| return cls( | |
| hidden_size, | |
| OURO_TRANSITION_COUNT, | |
| rank=rank, | |
| rmsnorm_epsilon=rmsnorm_epsilon, | |
| ) | |
| def for_huginn( | |
| cls, | |
| hidden_size: int, | |
| *, | |
| rank: int = DEFAULT_CTA_RANK, | |
| rmsnorm_epsilon: float = DEFAULT_RMSNORM_EPSILON, | |
| ) -> "CrossLoopTransitionAdapter": | |
| return cls( | |
| hidden_size, | |
| HUGINN_TRANSITION_COUNT, | |
| rank=rank, | |
| rmsnorm_epsilon=rmsnorm_epsilon, | |
| ) | |
| def _rmsnorm(self, hidden_state: torch.Tensor) -> torch.Tensor: | |
| work = hidden_state.to(torch.float32) | |
| normalized = work * torch.rsqrt(work.square().mean(dim=-1, keepdim=True) + self.rmsnorm_epsilon) | |
| return normalized.to(hidden_state.dtype) | |
| def forward(self, hidden_state: torch.Tensor, transition_index: int) -> torch.Tensor: | |
| if hidden_state.ndim == 0 or hidden_state.shape[-1] != self.hidden_size: | |
| raise ValueError( | |
| f"hidden_state last dimension must be {self.hidden_size}, " | |
| f"got {tuple(hidden_state.shape)}" | |
| ) | |
| if not 0 <= transition_index < self.transition_count: | |
| raise IndexError( | |
| f"transition_index must be in [0, {self.transition_count}), " | |
| f"got {transition_index}" | |
| ) | |
| if not self.enabled: | |
| return hidden_state | |
| normalized = self._rmsnorm(hidden_state) | |
| a = self.a[transition_index].to(hidden_state) | |
| b = self.b[transition_index].to(hidden_state) | |
| eta = self.eta[transition_index].to(hidden_state) | |
| U = self.U.to(hidden_state) | |
| V = self.V.to(hidden_state) | |
| affine = (a - 1) * normalized + b | |
| low_rank = ((normalized @ V) * eta) @ U.T | |
| return hidden_state + affine + low_rank | |
| def export_state(self) -> dict[str, Any]: | |
| return { | |
| "format": "loopq_cta", | |
| "format_version": self.FORMAT_VERSION, | |
| "hidden_size": self.hidden_size, | |
| "transition_count": self.transition_count, | |
| "rank": self.rank, | |
| "enabled": self.enabled, | |
| "rmsnorm_epsilon": self.rmsnorm_epsilon, | |
| "state_dict": {name: value.detach().cpu() for name, value in self.state_dict().items()}, | |
| } | |
| def from_export_state(cls, state: Mapping[str, Any]) -> "CrossLoopTransitionAdapter": | |
| required = { | |
| "format", | |
| "format_version", | |
| "hidden_size", | |
| "transition_count", | |
| "rank", | |
| "rmsnorm_epsilon", | |
| "state_dict", | |
| } | |
| missing = required.difference(state) | |
| if missing: | |
| raise ValueError(f"CTA export is missing fields: {sorted(missing)}") | |
| if state["format"] != "loopq_cta" or state["format_version"] != cls.FORMAT_VERSION: | |
| raise ValueError("unsupported CTA export format or version") | |
| adapter = cls( | |
| int(state["hidden_size"]), | |
| int(state["transition_count"]), | |
| rank=int(state["rank"]), | |
| rmsnorm_epsilon=float(state["rmsnorm_epsilon"]), | |
| dtype=state["state_dict"]["U"].dtype, | |
| ) | |
| adapter.load_state_dict(state["state_dict"]) | |
| adapter.enabled = bool(state.get("enabled", True)) | |
| if not adapter.enabled: | |
| adapter.requires_grad_(False) | |
| return adapter | |