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