Download loopq_quantization/scripts/loopq/transforms.py from JunYoungLee/ut-depth-probe-artifacts: direct link, hf CLI and curl.
- Browser
- Download file 11.8 kB
-
https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/transforms.py
- Command line
-
hf download hf://JunYoungLee/ut-depth-probe-artifacts/loopq_quantization/scripts/loopq/transforms.py
-
curl -L -o transforms.py https://huggingface.co/JunYoungLee/ut-depth-probe-artifacts/resolve/main/loopq_quantization/scripts/loopq/transforms.py
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) | |
| 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)) | |
| 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()}, | |
| } | |
| 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)) | |
| 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 | |
| ) | |
| def feature_size(self) -> int: | |
| return self.left_size * self.right_size | |
| def left(self) -> torch.Tensor: | |
| return self.left_u.weight @ torch.diag(self.left_singular) @ self.left_v.weight.T | |
| def right(self) -> torch.Tensor: | |
| return self.right_u.weight @ torch.diag(self.right_singular) @ self.right_v.weight.T | |
| 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(), | |
| }, | |
| } | |