JunYoungLee's picture
Add LoopQ 4-bit quantization of Ouro-1.4B
9118991 verified
Raw History Blame Contribute Delete
11.8 kB
"""Shared Kronecker transform and exact no-quant folding for LoopQ LQ3.
LoopQ Equations (2)-(3) require ``X @ P`` to be paired with
``W @ P^{-T}``. The paper additionally says that ``P`` uses a
FlatQuant-style Kronecker decomposition. This module implements that minimal
shared contract without choosing LoopQ's underspecified optimizer or SVD/direct
calibration parameterization.
"""
from __future__ import annotations
import torch
from torch import nn
from torch.nn.utils import parametrizations
class MaterializedKroneckerTransform:
# One-forward execution view reusing effective factors and inverses.
def __init__(self, transform: nn.Module) -> None:
self.feature_size = transform.feature_size
self.left = transform.left
self.right = transform.right
log_diagonal = transform.log_diagonal
self.diagonal = None if log_diagonal is None else log_diagonal.exp()
self.left_inverse_transpose = torch.linalg.inv(
self.left.to(torch.float64)
).T.to(dtype=self.left.dtype)
self.right_inverse_transpose = torch.linalg.inv(
self.right.to(torch.float64)
).T.to(dtype=self.right.dtype)
def __call__(self, activation: torch.Tensor) -> torch.Tensor:
if activation.ndim == 0 or activation.shape[-1] != self.feature_size:
raise ValueError(
f"activation last dimension must be {self.feature_size}, got {tuple(activation.shape)}"
)
value = activation
if self.diagonal is not None:
value = value * self.diagonal.to(value)
return SharedKroneckerTransform._apply_kronecker(
value, self.left.to(value), self.right.to(value)
)
def fold_weight(self, weight: torch.Tensor) -> torch.Tensor:
if weight.ndim != 2 or weight.shape[1] != self.feature_size:
raise ValueError(
f"weight must have shape (out_features, {self.feature_size}), got {tuple(weight.shape)}"
)
value = weight
if self.diagonal is not None:
value = value / self.diagonal.to(value)
return SharedKroneckerTransform._apply_kronecker(
value,
self.left_inverse_transpose.to(value),
self.right_inverse_transpose.to(value),
)
class SharedKroneckerTransform(nn.Module):
"""One loop-shared invertible transform ``diag(d) @ kron(L, R)``.
The explicit factors make architecture-selected factor dimensions part of
the artifact rather than silently guessing them from hidden size. Identity
initialization exactly preserves the original BF16 function.
"""
def __init__(
self,
left_size: int,
right_size: int,
*,
add_diagonal: bool = False,
dtype: torch.dtype = torch.float32,
) -> None:
super().__init__()
if left_size <= 0 or right_size <= 0:
raise ValueError("Kronecker factor sizes must be positive")
self.left_size = int(left_size)
self.right_size = int(right_size)
self.add_diagonal = bool(add_diagonal)
self.left = nn.Parameter(torch.eye(left_size, dtype=dtype))
self.right = nn.Parameter(torch.eye(right_size, dtype=dtype))
if add_diagonal:
self.log_diagonal = nn.Parameter(torch.zeros(self.feature_size, dtype=dtype))
else:
self.register_parameter("log_diagonal", None)
@property
def feature_size(self) -> int:
return self.left_size * self.right_size
def matrix(self) -> torch.Tensor:
"""Materialize ``P`` for correctness paths and export."""
kronecker = torch.kron(self.left, self.right)
if self.log_diagonal is None:
return kronecker
return self.log_diagonal.exp().diag() @ kronecker
def forward(self, activation: torch.Tensor) -> torch.Tensor:
if activation.ndim == 0 or activation.shape[-1] != self.feature_size:
raise ValueError(
f"activation last dimension must be {self.feature_size}, "
f"got {tuple(activation.shape)}"
)
value = activation
if self.log_diagonal is not None:
value = value * self.log_diagonal.exp().to(value)
return self._apply_kronecker(value, self.left.to(value), self.right.to(value))
@staticmethod
def _apply_kronecker(
value: torch.Tensor, left: torch.Tensor, right: torch.Tensor
) -> torch.Tensor:
"""Compute ``value @ kron(left, right)`` without materializing kron."""
shape = value.shape
reshaped = value.reshape(-1, left.shape[0], right.shape[0])
reshaped = torch.matmul(reshaped, right)
reshaped = torch.matmul(left.T, reshaped)
return reshaped.reshape(shape)
def materialize(self) -> MaterializedKroneckerTransform:
return MaterializedKroneckerTransform(self)
def inverse_transpose(self) -> torch.Tensor:
"""Return ``P^{-T}``, computed in FP64 for stable offline folding."""
matrix = self.matrix()
inverse_transpose = torch.linalg.inv(matrix.to(torch.float64)).T
return inverse_transpose.to(dtype=matrix.dtype)
def fold_weight(self, weight: torch.Tensor) -> torch.Tensor:
"""Return ``W @ P^{-T}`` without mutating the supplied shared weight."""
if weight.ndim != 2 or weight.shape[1] != self.feature_size:
raise ValueError(
f"weight must have shape (out_features, {self.feature_size}), "
f"got {tuple(weight.shape)}"
)
value = weight
if self.log_diagonal is not None:
value = value / self.log_diagonal.exp().to(value)
left_inverse_transpose = torch.linalg.inv(self.left.to(torch.float64)).T.to(value)
right_inverse_transpose = torch.linalg.inv(self.right.to(torch.float64)).T.to(value)
return self._apply_kronecker(value, left_inverse_transpose, right_inverse_transpose)
def fresh_copy(self) -> "SharedKroneckerTransform":
result = type(self)(
self.left_size, self.right_size, add_diagonal=self.add_diagonal,
dtype=self.left.dtype,
).to(device=self.left.device)
result.load_state_dict(self.state_dict())
return result
def export_state(self) -> dict[str, object]:
return {
"format": "loopq_shared_kronecker_transform",
"format_version": 1,
"left_size": self.left_size,
"right_size": self.right_size,
"add_diagonal": self.add_diagonal,
"state_dict": {key: value.detach().cpu() for key, value in self.state_dict().items()},
}
@classmethod
def from_export_state(cls, state: dict[str, object]) -> "SharedKroneckerTransform":
if state.get("format") != "loopq_shared_kronecker_transform" or state.get("format_version") != 1:
raise ValueError("unsupported shared-transform export format or version")
transform = cls(
int(state["left_size"]),
int(state["right_size"]),
add_diagonal=bool(state["add_diagonal"]),
dtype=state["state_dict"]["left"].dtype,
)
transform.load_state_dict(state["state_dict"])
return transform
def _random_orthogonal(size: int, *, dtype: torch.dtype) -> torch.Tensor:
"""Return a torch-RNG-controlled analogue of FlatQuant get_init_weight."""
value = torch.randn(size, size, dtype=dtype)
q, r = torch.linalg.qr(value)
signs = torch.where(torch.diagonal(r) < 0, -torch.ones((), dtype=dtype),
torch.ones((), dtype=dtype))
return q @ torch.diag(signs)
class FlatQuantSVDKroneckerTransform(nn.Module):
"""FlatQuant-default SVD parameterization of a Kronecker transform.
Each small Kronecker factor is ``U diag(s) V.T``. ``U`` and ``V`` use
PyTorch's Cayley orthogonal parameterization, matching the cited official
FlatQuant implementation. Export stores only the effective factors so
inference does not depend on the training parameterization.
"""
def __init__(self, left_size: int, right_size: int, *,
dtype: torch.dtype = torch.float32) -> None:
super().__init__()
if left_size <= 0 or right_size <= 0:
raise ValueError("Kronecker factor sizes must be positive")
self.left_size = int(left_size)
self.right_size = int(right_size)
self.add_diagonal = False
self.left_u = self._orthogonal_linear(left_size, dtype)
self.left_v = self._orthogonal_linear(left_size, dtype)
self.right_u = self._orthogonal_linear(right_size, dtype)
self.right_v = self._orthogonal_linear(right_size, dtype)
self.left_singular = nn.Parameter(torch.ones(left_size, dtype=dtype))
self.right_singular = nn.Parameter(torch.ones(right_size, dtype=dtype))
@staticmethod
def _orthogonal_linear(size: int, dtype: torch.dtype) -> nn.Module:
linear = nn.Linear(size, size, bias=False, dtype=dtype)
with torch.no_grad():
linear.weight.copy_(_random_orthogonal(size, dtype=dtype))
return parametrizations.orthogonal(
linear, orthogonal_map="cayley", use_trivialization=False
)
@property
def feature_size(self) -> int:
return self.left_size * self.right_size
@property
def left(self) -> torch.Tensor:
return self.left_u.weight @ torch.diag(self.left_singular) @ self.left_v.weight.T
@property
def right(self) -> torch.Tensor:
return self.right_u.weight @ torch.diag(self.right_singular) @ self.right_v.weight.T
@property
def log_diagonal(self):
return None
def matrix(self) -> torch.Tensor:
return torch.kron(self.left, self.right)
def materialize(self) -> MaterializedKroneckerTransform:
return MaterializedKroneckerTransform(self)
def forward(self, activation: torch.Tensor) -> torch.Tensor:
if activation.ndim == 0 or activation.shape[-1] != self.feature_size:
raise ValueError(
f"activation last dimension must be {self.feature_size}, got {tuple(activation.shape)}"
)
return SharedKroneckerTransform._apply_kronecker(
activation, self.left.to(activation), self.right.to(activation)
)
def fold_weight(self, weight: torch.Tensor) -> torch.Tensor:
if weight.ndim != 2 or weight.shape[1] != self.feature_size:
raise ValueError(
f"weight must have shape (out_features, {self.feature_size}), got {tuple(weight.shape)}"
)
left_inverse_transpose = torch.linalg.inv(self.left.to(torch.float64)).T.to(weight)
right_inverse_transpose = torch.linalg.inv(self.right.to(torch.float64)).T.to(weight)
return SharedKroneckerTransform._apply_kronecker(
weight, left_inverse_transpose, right_inverse_transpose
)
def fresh_copy(self) -> "FlatQuantSVDKroneckerTransform":
result = type(self)(self.left_size, self.right_size).to(
device=self.left_singular.device, dtype=self.left_singular.dtype
)
result.load_state_dict(self.state_dict())
return result
def export_state(self) -> dict[str, object]:
return {
"format": "loopq_shared_kronecker_transform",
"format_version": 1,
"left_size": self.left_size,
"right_size": self.right_size,
"add_diagonal": False,
"parameterization": "flatquant_svd_cayley",
"state_dict": {
"left": self.left.detach().cpu(),
"right": self.right.detach().cpu(),
},
}