"""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(), }, }