Added experimental orbit quant
This commit is contained in:
@@ -30,6 +30,7 @@ LINEAR_MODULES = [
|
||||
'Linear',
|
||||
'LoRACompatibleLinear',
|
||||
'QLinear',
|
||||
'OstrisLinear',
|
||||
# 'GroupNorm',
|
||||
]
|
||||
CONV_MODULES = [
|
||||
|
||||
@@ -29,7 +29,8 @@ Module = Union['LoConSpecialModule', 'LoRAModule', 'DoRAModule']
|
||||
LINEAR_MODULES = [
|
||||
'Linear',
|
||||
'LoRACompatibleLinear',
|
||||
'QLinear'
|
||||
'QLinear',
|
||||
'OstrisLinear',
|
||||
# 'GroupNorm',
|
||||
]
|
||||
CONV_MODULES = [
|
||||
|
||||
226
toolkit/util/orbit_quant.py
Normal file
226
toolkit/util/orbit_quant.py
Normal file
@@ -0,0 +1,226 @@
|
||||
"""
|
||||
OrbitQuant weight-only quantization backend (orbit2 / orbit3 / orbit4 qtypes).
|
||||
|
||||
Implements the weight half of "OrbitQuant: Data-Agnostic Quantization for Image and
|
||||
Video Diffusion Transformers" (arXiv:2607.02461) as an OstrisQuantizer backend (see
|
||||
toolkit/util/ostris_quant.py). Each linear weight is rotated with a randomized
|
||||
permuted block-Hadamard (RPBH) rotation shared per input dimension, split into
|
||||
per-row norms and unit directions, and the directions are quantized with a Lloyd-Max
|
||||
codebook fit to the fixed post-rotation coordinate marginal N(0, 1/d).
|
||||
No calibration data is needed.
|
||||
|
||||
At runtime the forward rotation is applied to the activations instead of un-rotating
|
||||
the weight, so the two cancel in the matmul (rotate = multiply by P):
|
||||
|
||||
W' = W P^T, y = dequant(W') (P x) ~= W x
|
||||
|
||||
Activations are not quantized. This is meant for holding a frozen base model at low
|
||||
bit-width while training adapters on top; activation quantization would only add
|
||||
noise without fused low-bit kernels.
|
||||
|
||||
Quantized state attached to each module:
|
||||
orbit_packed packed codebook indices (uint8 bitstream)
|
||||
orbit_row_norms per output row l2 norms of the rotated weight (original dtype)
|
||||
orbit_codebook Lloyd-Max centroids for N(0, 1/in_features) (float32)
|
||||
orbit_perm / orbit_inv_perm / orbit_signs shared RPBH rotation for in_features
|
||||
orbit_bits / orbit_block bit-width and Hadamard block size
|
||||
"""
|
||||
|
||||
import math
|
||||
from typing import Dict, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from toolkit.print import print_acc
|
||||
from toolkit.util.ostris_quant import OstrisQuantizer
|
||||
|
||||
# qtype name -> bits per weight coordinate
|
||||
ORBIT_QTYPES = {"orbit2": 2, "orbit3": 3, "orbit4": 4}
|
||||
|
||||
# below this Hadamard block size the rotated coordinates are not gaussian enough for
|
||||
# the shared N(0, 1/d) codebook to be valid, so conversion is skipped for the layer
|
||||
MIN_HADAMARD_BLOCK = 32
|
||||
|
||||
_normal_codebook_cache: Dict[int, torch.Tensor] = {}
|
||||
_rotation_cache: Dict[int, Tuple[torch.Tensor, torch.Tensor]] = {}
|
||||
_skip_warned = set()
|
||||
|
||||
|
||||
def gaussian_lloyd_max(bits: int, iters: int = 200) -> torch.Tensor:
|
||||
"""MSE-optimal (Lloyd-Max) centroids for the standard normal, 2**bits levels,
|
||||
returned in ascending order as float32. Cached per bit-width."""
|
||||
if bits in _normal_codebook_cache:
|
||||
return _normal_codebook_cache[bits]
|
||||
levels = 2 ** bits
|
||||
# init at the gaussian quantile midpoints
|
||||
q = (torch.arange(levels, dtype=torch.float64) + 0.5) / levels
|
||||
c = math.sqrt(2.0) * torch.erfinv(2.0 * q - 1.0)
|
||||
inf = torch.tensor([math.inf], dtype=torch.float64)
|
||||
for _ in range(iters):
|
||||
edges = (c[:-1] + c[1:]) / 2.0
|
||||
lo = torch.cat([-inf, edges])
|
||||
hi = torch.cat([edges, inf])
|
||||
# centroid update: E[X | lo < X < hi] = (phi(lo) - phi(hi)) / (Phi(hi) - Phi(lo))
|
||||
phi_lo = torch.exp(-lo * lo / 2.0) / math.sqrt(2.0 * math.pi)
|
||||
phi_hi = torch.exp(-hi * hi / 2.0) / math.sqrt(2.0 * math.pi)
|
||||
cdf_lo = 0.5 * (1.0 + torch.erf(lo / math.sqrt(2.0)))
|
||||
cdf_hi = 0.5 * (1.0 + torch.erf(hi / math.sqrt(2.0)))
|
||||
c = (phi_lo - phi_hi) / (cdf_hi - cdf_lo)
|
||||
c = c.to(torch.float32)
|
||||
_normal_codebook_cache[bits] = c
|
||||
return c
|
||||
|
||||
|
||||
def hadamard_block_size(d: int) -> int:
|
||||
# largest power of two dividing d
|
||||
return d & (-d)
|
||||
|
||||
|
||||
def rpbh_params(d: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Permutation (int64) and Rademacher signs (int8) of the RPBH rotation for input
|
||||
dimension d, on cpu. Sampled once per dimension with a seed derived from d so the
|
||||
rotation is identical across layers, runs, and resumes."""
|
||||
if d in _rotation_cache:
|
||||
return _rotation_cache[d]
|
||||
g = torch.Generator().manual_seed(0x0EB17 + d)
|
||||
perm = torch.randperm(d, generator=g)
|
||||
signs = torch.randint(0, 2, (d,), generator=g, dtype=torch.int8) * 2 - 1
|
||||
_rotation_cache[d] = (perm, signs)
|
||||
return perm, signs
|
||||
|
||||
|
||||
def _fwht(x: torch.Tensor, h: int) -> torch.Tensor:
|
||||
"""Orthonormal fast Walsh-Hadamard transform applied to each contiguous block of h
|
||||
coordinates along the last dimension (h must be a power of two dividing the dim)."""
|
||||
shape = x.shape
|
||||
x = x.reshape(-1, h)
|
||||
m = x.shape[0]
|
||||
step = 1
|
||||
while step < h:
|
||||
y = x.view(m, h // (2 * step), 2, step)
|
||||
x = torch.stack((y[:, :, 0] + y[:, :, 1], y[:, :, 0] - y[:, :, 1]), dim=2).view(m, h)
|
||||
step *= 2
|
||||
return (x * h ** -0.5).view(shape)
|
||||
|
||||
|
||||
def rpbh_forward(x: torch.Tensor, perm: torch.Tensor, signs: torch.Tensor, h: int) -> torch.Tensor:
|
||||
"""y = blkdiag(H_h D) P x applied to the last dimension of x."""
|
||||
y = torch.index_select(x, -1, perm) * signs.to(x.dtype)
|
||||
return _fwht(y, h)
|
||||
|
||||
|
||||
def rpbh_inverse(y: torch.Tensor, inv_perm: torch.Tensor, signs: torch.Tensor, h: int) -> torch.Tensor:
|
||||
"""Inverse of rpbh_forward (the rotation is orthogonal, H is self-inverse)."""
|
||||
z = _fwht(y, h) * signs.to(y.dtype)
|
||||
return torch.index_select(z, -1, inv_perm)
|
||||
|
||||
|
||||
def pack_codes(codes: torch.Tensor, bits: int) -> torch.Tensor:
|
||||
"""Pack integer codes (values < 2**bits, any shape) into a flat uint8 bitstream."""
|
||||
flat = codes.flatten().to(torch.uint8)
|
||||
pad = (-flat.numel()) % 8
|
||||
if pad:
|
||||
flat = torch.cat([flat, flat.new_zeros(pad)])
|
||||
shifts = torch.arange(bits - 1, -1, -1, device=flat.device, dtype=torch.uint8)
|
||||
bit_mat = (flat.unsqueeze(-1) >> shifts) & 1 # (n, bits)
|
||||
byte_mat = bit_mat.view(-1, 8)
|
||||
weights = torch.tensor([1 << i for i in range(7, -1, -1)], device=flat.device, dtype=torch.uint8)
|
||||
return (byte_mat * weights).sum(-1, dtype=torch.uint8)
|
||||
|
||||
|
||||
def unpack_codes(packed: torch.Tensor, bits: int, numel: int) -> torch.Tensor:
|
||||
"""Inverse of pack_codes. Returns a flat uint8 tensor of length numel."""
|
||||
shifts = torch.arange(7, -1, -1, device=packed.device, dtype=torch.uint8)
|
||||
bit_mat = ((packed.unsqueeze(-1) >> shifts) & 1).view(-1, bits)
|
||||
weights = torch.tensor([1 << i for i in range(bits - 1, -1, -1)], device=packed.device, dtype=torch.uint8)
|
||||
codes = (bit_mat * weights).sum(-1, dtype=torch.uint8)
|
||||
return codes[:numel]
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def _quantize_rows(
|
||||
w_fp32: torch.Tensor,
|
||||
perm: torch.Tensor,
|
||||
signs: torch.Tensor,
|
||||
h: int,
|
||||
codebook: torch.Tensor,
|
||||
bits: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Rotate weight rows into the RPBH basis and quantize their unit directions
|
||||
against the codebook. Returns (packed codes, float32 row norms)."""
|
||||
w_rot = rpbh_forward(w_fp32, perm, signs, h)
|
||||
row_norms = w_rot.norm(dim=1)
|
||||
unit = w_rot / (row_norms + 1e-10).unsqueeze(1)
|
||||
edges = (codebook[:-1] + codebook[1:]) / 2
|
||||
codes = torch.bucketize(unit, edges, out_int32=True).to(torch.uint8)
|
||||
return pack_codes(codes, bits), row_norms
|
||||
|
||||
|
||||
class OrbitQuantizer(OstrisQuantizer):
|
||||
"""OrbitQuant backend. One instance per bit-width, shareable across modules."""
|
||||
|
||||
def __init__(self, bits: int):
|
||||
self.bits = bits
|
||||
|
||||
def can_quantize(self, module: torch.nn.Linear) -> bool:
|
||||
d = module.in_features
|
||||
h = hadamard_block_size(d)
|
||||
if h < MIN_HADAMARD_BLOCK:
|
||||
if d not in _skip_warned:
|
||||
_skip_warned.add(d)
|
||||
print_acc(
|
||||
f"OrbitQuant: skipping linears with in_features={d} "
|
||||
f"(power-of-two block {h} is too small for the rotation)"
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
def quantize_(self, module: torch.nn.Linear, weight_fp32: torch.Tensor) -> None:
|
||||
d = module.in_features
|
||||
h = hadamard_block_size(d)
|
||||
device = weight_fp32.device
|
||||
perm_cpu, signs_cpu = rpbh_params(d)
|
||||
perm = perm_cpu.to(device=device, dtype=torch.int32)
|
||||
inv_perm = torch.argsort(perm_cpu).to(device=device, dtype=torch.int32)
|
||||
signs = signs_cpu.to(device)
|
||||
codebook = (gaussian_lloyd_max(self.bits) * d ** -0.5).to(device)
|
||||
packed, row_norms = _quantize_rows(weight_fp32, perm, signs, h, codebook, self.bits)
|
||||
module.register_buffer("orbit_packed", packed, persistent=False)
|
||||
module.register_buffer("orbit_row_norms", row_norms.to(module.weight.dtype), persistent=False)
|
||||
module.register_buffer("orbit_codebook", codebook, persistent=False)
|
||||
module.register_buffer("orbit_perm", perm, persistent=False)
|
||||
module.register_buffer("orbit_inv_perm", inv_perm, persistent=False)
|
||||
module.register_buffer("orbit_signs", signs, persistent=False)
|
||||
module.orbit_bits = self.bits
|
||||
module.orbit_block = h
|
||||
|
||||
def _dequantize_rotated(self, module, dtype: torch.dtype) -> torch.Tensor:
|
||||
"""Materialize the rotated-basis weight W' = W P^T in the given dtype."""
|
||||
numel = module.out_features * module.in_features
|
||||
codes = unpack_codes(module.orbit_packed, module.orbit_bits, numel)
|
||||
w = torch.index_select(module.orbit_codebook.to(dtype), 0, codes.to(torch.int32))
|
||||
w = w.view(module.out_features, module.in_features)
|
||||
return w * module.orbit_row_norms.to(dtype).unsqueeze(1)
|
||||
|
||||
def dequantize(self, module) -> torch.Tensor:
|
||||
w = self._dequantize_rotated(module, torch.float32)
|
||||
return rpbh_inverse(w, module.orbit_inv_perm, module.orbit_signs, module.orbit_block)
|
||||
|
||||
def requantize_(self, module, fp_weight: torch.Tensor) -> None:
|
||||
w = fp_weight.to(device=module.orbit_packed.device, dtype=torch.float32)
|
||||
packed, row_norms = _quantize_rows(
|
||||
w, module.orbit_perm, module.orbit_signs, module.orbit_block,
|
||||
module.orbit_codebook, module.orbit_bits,
|
||||
)
|
||||
module.orbit_packed = packed
|
||||
module.orbit_row_norms = row_norms.to(module.ostris_orig_dtype)
|
||||
|
||||
def forward(self, module, x: torch.Tensor) -> torch.Tensor:
|
||||
# rotate the activation instead of un-rotating the weight; the rotations
|
||||
# cancel in the matmul. the weight is frozen, so build it outside autograd;
|
||||
# gradients still flow to x through the rotation and the matmul
|
||||
with torch.no_grad():
|
||||
w = self._dequantize_rotated(module, x.dtype)
|
||||
x_rot = rpbh_forward(x, module.orbit_perm, module.orbit_signs, module.orbit_block)
|
||||
return F.linear(x_rot, w, module.bias)
|
||||
365
toolkit/util/orbit_vq_quant.py
Normal file
365
toolkit/util/orbit_vq_quant.py
Normal file
@@ -0,0 +1,365 @@
|
||||
"""
|
||||
OrbitVQ weight-only quantization backend (orbitvq2 / orbitvq3 / orbitvq4 qtypes).
|
||||
|
||||
Extends the OrbitQuant recipe (toolkit/util/orbit_quant.py) with three accuracy
|
||||
upgrades aimed at low bit-widths, at the cost of no longer being the paper's pure
|
||||
scalar method:
|
||||
|
||||
1. Lattice vector codebooks. Groups of coordinates are quantized jointly against a
|
||||
codebook built from the densest lattice for the group dimension (E8 for 8-dim
|
||||
groups at 2 bits, D4 for 4-dim groups at 3/4 bits) instead of one coordinate at
|
||||
a time. A scalar codebook wastes indices on corner combinations that iid
|
||||
gaussian coordinates never produce; a lattice codebook spends all its codewords
|
||||
on the spherical shell where rotated weights actually live.
|
||||
2. Per-group scales. One scale per GROUP_SIZE rotated coordinates (instead of one
|
||||
norm per row) absorbs the finite-sample energy fluctuation between groups.
|
||||
3. Least-squares scale refit. After codes are chosen the scale is refit to the
|
||||
MSE-optimal value <w, c> / <c, c>.
|
||||
|
||||
Everything stays data-free and deterministic: the same seeded RPBH rotation as
|
||||
OrbitQuant, codebooks that are fixed mathematical objects (lattice points sorted by
|
||||
norm), and precomputed distortion-optimal lattice scales for the N(0,1) source.
|
||||
|
||||
Encoding uses the closed-form nearest-lattice-point algorithm plus a hash lookup
|
||||
into the truncated codebook; the rare vectors whose nearest lattice point falls
|
||||
outside the codebook fall back to an exact brute-force search.
|
||||
|
||||
Bits per param: index bits exactly (2/3/4) + 16-bit group scales / GROUP_SIZE
|
||||
(+0.125 at the default 128).
|
||||
|
||||
Quantized state attached to each module:
|
||||
ovq_packed packed codeword indices (uint8 bitstream, index_bits per vector)
|
||||
ovq_scales per (row, group) least-squares scales (original dtype)
|
||||
ovq_perm / ovq_inv_perm / ovq_signs shared RPBH rotation for in_features
|
||||
ovq_block / ovq_group Hadamard block size and group size
|
||||
"""
|
||||
|
||||
import math
|
||||
from typing import Dict, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from toolkit.print import print_acc
|
||||
from toolkit.util.ostris_quant import OstrisQuantizer
|
||||
from toolkit.util.orbit_quant import (
|
||||
MIN_HADAMARD_BLOCK,
|
||||
hadamard_block_size,
|
||||
rpbh_forward,
|
||||
rpbh_inverse,
|
||||
rpbh_params,
|
||||
)
|
||||
|
||||
# qtype name -> quantizer constructor kwargs
|
||||
ORBIT_VQ_QTYPES = {
|
||||
"orbitvq2": {"bits": 2, "vec_dim": 8, "lattice": "E8", "codebook_size": 2 ** 16},
|
||||
"orbitvq3": {"bits": 3, "vec_dim": 4, "lattice": "D4", "codebook_size": 2 ** 12},
|
||||
"orbitvq4": {"bits": 4, "vec_dim": 4, "lattice": "D4", "codebook_size": 2 ** 16},
|
||||
}
|
||||
|
||||
# coordinates per stored scale (the group-norm granularity)
|
||||
GROUP_SIZE = 128
|
||||
|
||||
# encode -> least-squares scale refit rounds. re-encoding with the refit scale
|
||||
# (rounds=2) measured only 0.0005 lower relative error at 2-bit while doubling the
|
||||
# encode cost, so one round is the default
|
||||
LS_REFIT_ROUNDS = 1
|
||||
|
||||
# distortion-optimal lattice scale for a unit-variance gaussian source, per
|
||||
# (lattice, codebook_size), found by an exact nearest-codeword sweep over a seeded
|
||||
# 60k gaussian sample. per-coordinate MSE at these scales (vs scalar Lloyd-Max):
|
||||
# E8-65536 2-bit: 0.0916 (0.1175) D4-4096 3-bit: 0.0301 (0.0345)
|
||||
# D4-65536 4-bit: 0.0089 (0.0095)
|
||||
# fixed constants so codebooks are bit-identical everywhere.
|
||||
BETA = {
|
||||
("E8", 2 ** 16): 0.9800,
|
||||
("D4", 2 ** 12): 0.4722,
|
||||
("D4", 2 ** 16): 0.2617,
|
||||
}
|
||||
|
||||
# key packing for the hash lookup: doubled coordinates + _KEY_OFFSET must fit in
|
||||
# _KEY_BITS bits per dimension (covers every codebook point of the sizes above)
|
||||
_KEY_BITS = 6
|
||||
_KEY_OFFSET = 32
|
||||
|
||||
_skip_warned = set()
|
||||
_master_tables: Dict[Tuple[str, int], "_VQTables"] = {}
|
||||
_device_tables: Dict[Tuple[str, int, str], "_VQTables"] = {}
|
||||
|
||||
|
||||
def enumerate_lattice_codebook(lattice: str, size: int) -> torch.Tensor:
|
||||
"""The `size` lattice points closest to the origin (ties broken lexicographically,
|
||||
so the result is fully deterministic), as float32 (size, dim), in lattice units.
|
||||
|
||||
D4 = {x in Z^4 : sum(x) even} (densest 4-dim packing).
|
||||
E8 = D8 union (D8 + 1/2) (densest 8-dim packing). In doubled coordinates both
|
||||
become: uniform-parity integer vectors with sum divisible by 4.
|
||||
"""
|
||||
if lattice == "D4":
|
||||
dim, reach = 4, 27 # doubled coords in [-26, 26], covers 65536 points
|
||||
vals = torch.arange(-(reach - 1), reach, 2, dtype=torch.int32) # even only
|
||||
parities = [vals]
|
||||
elif lattice == "E8":
|
||||
dim = 8
|
||||
# doubled coords: even in [-6, 6] (integer points, norm^2 <= 12) and odd in
|
||||
# [-5, 5] (half points); both ranges cover far more than 65536 points
|
||||
parities = [
|
||||
torch.arange(-6, 7, 2, dtype=torch.int32),
|
||||
torch.arange(-5, 6, 2, dtype=torch.int32),
|
||||
]
|
||||
else:
|
||||
raise ValueError(f"unknown lattice {lattice}")
|
||||
|
||||
kept = []
|
||||
for vals in parities:
|
||||
# chunk over the first coordinate to keep the cartesian product small
|
||||
for v0 in vals.tolist():
|
||||
rest = torch.cartesian_prod(*([vals] * (dim - 1)))
|
||||
pts = torch.cat(
|
||||
[torch.full((rest.shape[0], 1), v0, dtype=torch.int32), rest], dim=1
|
||||
)
|
||||
pts = pts[pts.sum(dim=1).remainder(4) == 0]
|
||||
norm2 = (pts.to(torch.int64) ** 2).sum(dim=1)
|
||||
# generous radius cut just to bound memory before the exact sort below
|
||||
keep = norm2 <= (48 if lattice == "E8" else 26 ** 2 + 1)
|
||||
kept.append(pts[keep])
|
||||
pts = torch.cat(kept)
|
||||
|
||||
# sort by (norm, lexicographic key) and take the closest `size` points
|
||||
norm2 = (pts.to(torch.int64) ** 2).sum(dim=1)
|
||||
key = _point_keys(pts)
|
||||
order = torch.argsort(norm2 * (1 << (_KEY_BITS * dim)) + key)
|
||||
pts = pts[order[:size]]
|
||||
if pts.shape[0] < size:
|
||||
raise RuntimeError(f"lattice enumeration too small for {lattice}/{size}")
|
||||
return pts.to(torch.float32) / 2.0 # back to lattice units
|
||||
|
||||
|
||||
def _point_keys(doubled_pts: torch.Tensor) -> torch.Tensor:
|
||||
"""int64 hash key per point from doubled integer coordinates."""
|
||||
digits = doubled_pts.to(torch.int64) + _KEY_OFFSET
|
||||
key = torch.zeros(doubled_pts.shape[0], dtype=torch.int64, device=doubled_pts.device)
|
||||
for i in range(doubled_pts.shape[1]):
|
||||
key = key | (digits[:, i] << (_KEY_BITS * i))
|
||||
return key
|
||||
|
||||
|
||||
def _round_Dn(x: torch.Tensor) -> torch.Tensor:
|
||||
"""Nearest point of D_n (integer vectors with even coordinate sum) to each row."""
|
||||
f = x.round()
|
||||
odd = f.to(torch.int64).sum(dim=-1).remainder(2) != 0
|
||||
err = x - f
|
||||
idx = err.abs().argmax(dim=-1, keepdim=True)
|
||||
step = torch.where(err.gather(-1, idx) >= 0, 1.0, -1.0).to(x.dtype)
|
||||
adjusted = f.scatter(-1, idx, f.gather(-1, idx) + step)
|
||||
return torch.where(odd.unsqueeze(-1), adjusted, f)
|
||||
|
||||
|
||||
def _round_lattice(x: torch.Tensor, lattice: str) -> torch.Tensor:
|
||||
"""Closed-form nearest lattice point (Conway & Sloane) to each row of x."""
|
||||
a = _round_Dn(x)
|
||||
if lattice == "D4":
|
||||
return a
|
||||
# E8 = D8 union (D8 + 1/2): take the closer of the two cosets
|
||||
b = _round_Dn(x - 0.5) + 0.5
|
||||
da = (x - a).square().sum(dim=-1)
|
||||
db = (x - b).square().sum(dim=-1)
|
||||
return torch.where((da <= db).unsqueeze(-1), a, b)
|
||||
|
||||
|
||||
class _VQTables:
|
||||
"""Codebook + hash tables for one (lattice, codebook_size), on one device."""
|
||||
|
||||
def __init__(self, lattice: str, size: int):
|
||||
self.lattice = lattice
|
||||
self.size = size
|
||||
self.beta = BETA[(lattice, size)]
|
||||
points = enumerate_lattice_codebook(lattice, size)
|
||||
self.codebook = points * self.beta # (size, dim) float32, source units
|
||||
keys = _point_keys((points * 2).to(torch.int32))
|
||||
self.sorted_keys, order = torch.sort(keys)
|
||||
self.key_to_index = order.to(torch.int32)
|
||||
# for the brute-force fallback: argmax(z.c - |c|^2/2) == nearest codeword
|
||||
self.half_sq_norms = self.codebook.square().sum(dim=1) / 2
|
||||
self.codebook_t = self.codebook.T.contiguous()
|
||||
|
||||
def to(self, device: torch.device) -> "_VQTables":
|
||||
out = object.__new__(_VQTables)
|
||||
out.lattice, out.size, out.beta = self.lattice, self.size, self.beta
|
||||
out.codebook = self.codebook.to(device)
|
||||
out.sorted_keys = self.sorted_keys.to(device)
|
||||
out.key_to_index = self.key_to_index.to(device)
|
||||
out.half_sq_norms = self.half_sq_norms.to(device)
|
||||
out.codebook_t = self.codebook_t.to(device)
|
||||
return out
|
||||
|
||||
|
||||
def get_vq_tables(lattice: str, size: int, device) -> _VQTables:
|
||||
mkey = (lattice, size)
|
||||
if mkey not in _master_tables:
|
||||
_master_tables[mkey] = _VQTables(lattice, size)
|
||||
dkey = (lattice, size, str(device))
|
||||
if dkey not in _device_tables:
|
||||
_device_tables[dkey] = _master_tables[mkey].to(torch.device(device))
|
||||
return _device_tables[dkey]
|
||||
|
||||
|
||||
def encode_vectors(z: torch.Tensor, tables: _VQTables) -> torch.Tensor:
|
||||
"""Exact nearest-codeword indices (int32) for rows of z (float32, source units).
|
||||
|
||||
Closed-form lattice rounding + hash lookup; rows whose nearest lattice point is
|
||||
outside the truncated codebook fall back to a brute-force search.
|
||||
"""
|
||||
dim = z.shape[-1]
|
||||
p = _round_lattice(z / tables.beta, tables.lattice)
|
||||
digits = (p * 2).round().to(torch.int64) + _KEY_OFFSET
|
||||
in_range = ((digits >= 0) & (digits < (1 << _KEY_BITS))).all(dim=-1)
|
||||
key = torch.zeros(z.shape[0], dtype=torch.int64, device=z.device)
|
||||
for i in range(dim):
|
||||
key = key | (digits[:, i].clamp(0, (1 << _KEY_BITS) - 1) << (_KEY_BITS * i))
|
||||
pos = torch.searchsorted(tables.sorted_keys, key).clamp(max=tables.size - 1)
|
||||
hit = in_range & (tables.sorted_keys.gather(0, pos) == key)
|
||||
idx = tables.key_to_index.gather(0, pos.to(torch.int64))
|
||||
|
||||
miss = ~hit
|
||||
n_miss = int(miss.sum())
|
||||
if n_miss > 0:
|
||||
# on cuda run the search in fp16 (tensor cores + half the score-matrix
|
||||
# traffic); the score gap between competing codewords for these overload
|
||||
# vectors is far above fp16 noise. cpu stays fp32.
|
||||
dt = torch.float16 if z.device.type == "cuda" else torch.float32
|
||||
z_miss = z[miss].to(dt)
|
||||
cb_t = tables.codebook_t.to(dt)
|
||||
half_norms = tables.half_sq_norms.to(dt)
|
||||
found = torch.empty(n_miss, dtype=torch.int32, device=z.device)
|
||||
# chunk rows so the score matrix stays ~256MB while each matmul stays large
|
||||
# enough to saturate the gpu (overload vectors are ~15-20% of the total)
|
||||
chunk = max(256, (2 ** 26) // tables.size)
|
||||
for s in range(0, n_miss, chunk):
|
||||
scores = z_miss[s:s + chunk] @ cb_t - half_norms
|
||||
found[s:s + chunk] = scores.argmax(dim=1).to(torch.int32)
|
||||
idx[miss] = found
|
||||
return idx
|
||||
|
||||
|
||||
def pack_indices(idx: torch.Tensor, bits: int) -> torch.Tensor:
|
||||
"""Pack integer indices (< 2**bits, bits <= 16) into a flat uint8 bitstream."""
|
||||
flat = idx.flatten().to(torch.int32)
|
||||
pad = (-flat.numel()) % 8
|
||||
if pad:
|
||||
flat = torch.cat([flat, flat.new_zeros(pad)])
|
||||
shifts = torch.arange(bits - 1, -1, -1, device=flat.device, dtype=torch.int32)
|
||||
bit_mat = ((flat.unsqueeze(-1) >> shifts) & 1).to(torch.uint8) # (n, bits)
|
||||
byte_mat = bit_mat.view(-1, 8)
|
||||
weights = torch.tensor([1 << i for i in range(7, -1, -1)], device=flat.device, dtype=torch.uint8)
|
||||
return (byte_mat * weights).sum(-1, dtype=torch.uint8)
|
||||
|
||||
|
||||
def unpack_indices(packed: torch.Tensor, bits: int, numel: int) -> torch.Tensor:
|
||||
"""Inverse of pack_indices. Returns a flat int32 tensor of length numel."""
|
||||
shifts = torch.arange(7, -1, -1, device=packed.device, dtype=torch.uint8)
|
||||
bit_mat = ((packed.unsqueeze(-1) >> shifts) & 1).view(-1, bits) # uint8
|
||||
idx = None
|
||||
# accumulate in <=8-bit chunks so intermediates stay uint8
|
||||
for c0 in range(0, bits, 8):
|
||||
cw = min(8, bits - c0)
|
||||
w = torch.tensor([1 << i for i in range(cw - 1, -1, -1)], device=packed.device, dtype=torch.uint8)
|
||||
part = (bit_mat[:, c0:c0 + cw] * w).sum(-1, dtype=torch.uint8).to(torch.int32)
|
||||
part = part << (bits - c0 - cw)
|
||||
idx = part if idx is None else idx | part
|
||||
return idx[:numel]
|
||||
|
||||
|
||||
class OrbitVQQuantizer(OstrisQuantizer):
|
||||
"""RPBH rotation + lattice vector codebook + per-group least-squares scales.
|
||||
One instance per qtype, shareable across modules."""
|
||||
|
||||
def __init__(self, bits: int, vec_dim: int, lattice: str, codebook_size: int,
|
||||
group_size: int = GROUP_SIZE):
|
||||
self.bits = bits
|
||||
self.vec_dim = vec_dim
|
||||
self.lattice = lattice
|
||||
self.codebook_size = codebook_size
|
||||
self.group_size = group_size
|
||||
self.index_bits = bits * vec_dim
|
||||
|
||||
def can_quantize(self, module: torch.nn.Linear) -> bool:
|
||||
d = module.in_features
|
||||
h = hadamard_block_size(d)
|
||||
if h < MIN_HADAMARD_BLOCK:
|
||||
if d not in _skip_warned:
|
||||
_skip_warned.add(d)
|
||||
print_acc(
|
||||
f"OrbitVQ: skipping linears with in_features={d} "
|
||||
f"(power-of-two block {h} is too small for the rotation)"
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
def _encode_rotated(self, w_rot: torch.Tensor, g: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Quantize a rotated weight (m, d) in float32. Returns (packed indices,
|
||||
float32 (m, d//g) group scales)."""
|
||||
tables = get_vq_tables(self.lattice, self.codebook_size, w_rot.device)
|
||||
m, d = w_rot.shape
|
||||
u = w_rot.view(m, d // g, g)
|
||||
# initial scale: per-group rms, so standardized coords are ~N(0, 1)
|
||||
scale = u.norm(dim=-1, keepdim=True) / g ** 0.5 + 1e-12
|
||||
idx = None
|
||||
for _ in range(LS_REFIT_ROUNDS):
|
||||
z = (u / scale).reshape(-1, self.vec_dim)
|
||||
idx = encode_vectors(z, tables)
|
||||
c = tables.codebook.index_select(0, idx).view(m, d // g, g)
|
||||
# least-squares optimal scale given the chosen codewords
|
||||
num = (u * c).sum(dim=-1, keepdim=True)
|
||||
den = c.square().sum(dim=-1, keepdim=True) + 1e-12
|
||||
scale = num / den
|
||||
return pack_indices(idx, self.index_bits), scale.view(m, d // g)
|
||||
|
||||
def quantize_(self, module: torch.nn.Linear, weight_fp32: torch.Tensor) -> None:
|
||||
d = module.in_features
|
||||
h = hadamard_block_size(d)
|
||||
g = min(self.group_size, h)
|
||||
device = weight_fp32.device
|
||||
perm_cpu, signs_cpu = rpbh_params(d)
|
||||
perm = perm_cpu.to(device=device, dtype=torch.int32)
|
||||
inv_perm = torch.argsort(perm_cpu).to(device=device, dtype=torch.int32)
|
||||
signs = signs_cpu.to(device)
|
||||
w_rot = rpbh_forward(weight_fp32, perm, signs, h)
|
||||
packed, scales = self._encode_rotated(w_rot, g)
|
||||
module.register_buffer("ovq_packed", packed, persistent=False)
|
||||
module.register_buffer("ovq_scales", scales.to(module.weight.dtype), persistent=False)
|
||||
module.register_buffer("ovq_perm", perm, persistent=False)
|
||||
module.register_buffer("ovq_inv_perm", inv_perm, persistent=False)
|
||||
module.register_buffer("ovq_signs", signs, persistent=False)
|
||||
module.ovq_block = h
|
||||
module.ovq_group = g
|
||||
|
||||
def _dequantize_rotated(self, module, dtype: torch.dtype) -> torch.Tensor:
|
||||
"""Materialize the rotated-basis weight W' = W P^T in the given dtype."""
|
||||
tables = get_vq_tables(self.lattice, self.codebook_size, module.ovq_packed.device)
|
||||
m, d = module.out_features, module.in_features
|
||||
g = module.ovq_group
|
||||
idx = unpack_indices(module.ovq_packed, self.index_bits, m * d // self.vec_dim)
|
||||
w = tables.codebook.to(dtype).index_select(0, idx).view(m, d // g, g)
|
||||
w = w * module.ovq_scales.to(dtype).unsqueeze(-1)
|
||||
return w.view(m, d)
|
||||
|
||||
def dequantize(self, module) -> torch.Tensor:
|
||||
w = self._dequantize_rotated(module, torch.float32)
|
||||
return rpbh_inverse(w, module.ovq_inv_perm, module.ovq_signs, module.ovq_block)
|
||||
|
||||
def requantize_(self, module, fp_weight: torch.Tensor) -> None:
|
||||
w = fp_weight.to(device=module.ovq_packed.device, dtype=torch.float32)
|
||||
w_rot = rpbh_forward(w, module.ovq_perm, module.ovq_signs, module.ovq_block)
|
||||
packed, scales = self._encode_rotated(w_rot, module.ovq_group)
|
||||
module.ovq_packed = packed
|
||||
module.ovq_scales = scales.to(module.ostris_orig_dtype)
|
||||
|
||||
def forward(self, module, x: torch.Tensor) -> torch.Tensor:
|
||||
# rotate the activation instead of un-rotating the weight; the rotations
|
||||
# cancel in the matmul. the weight is frozen, so build it outside autograd;
|
||||
# gradients still flow to x through the rotation and the matmul
|
||||
with torch.no_grad():
|
||||
w = self._dequantize_rotated(module, x.dtype)
|
||||
x_rot = rpbh_forward(x, module.ovq_perm, module.ovq_signs, module.ovq_block)
|
||||
return F.linear(x_rot, w, module.bias)
|
||||
133
toolkit/util/ostris_quant.py
Normal file
133
toolkit/util/ostris_quant.py
Normal file
@@ -0,0 +1,133 @@
|
||||
"""
|
||||
Quantization-agnostic custom quantized linear.
|
||||
|
||||
OstrisLinear is a drop-in nn.Linear replacement whose weight is held by a pluggable
|
||||
quantizer backend (OstrisQuantizer). Backends own the quantized representation
|
||||
(buffers + per-module attributes) and how the forward pass computes W x from it; the
|
||||
module and the rest of the toolkit stay backend agnostic. The first backend is
|
||||
OrbitQuant (toolkit/util/orbit_quant.py) via the orbit2/orbit3/orbit4 qtypes; add new
|
||||
backends by implementing OstrisQuantizer and resolving them in get_ostris_quantizer.
|
||||
|
||||
Modules are converted in place by convert_linear_to_ostris via class swap, so the
|
||||
original module object (and any references to it, e.g. LoRA org_module or parent
|
||||
containers) stays valid.
|
||||
"""
|
||||
|
||||
from typing import Dict, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class OstrisQuantizer:
|
||||
"""Base class for weight quantization backends used by OstrisLinear.
|
||||
|
||||
Backends are stateless with respect to tensors: everything tensor-shaped must be
|
||||
registered as a buffer on the module inside quantize_ (so device moves and dtype
|
||||
casts through nn.Module._apply keep working), and read back off the module in the
|
||||
other methods. One backend instance may be shared by many modules.
|
||||
"""
|
||||
|
||||
def can_quantize(self, module: torch.nn.Linear) -> bool:
|
||||
"""Whether this backend can quantize the given linear (e.g. shape constraints)."""
|
||||
return True
|
||||
|
||||
def quantize_(self, module: torch.nn.Linear, weight_fp32: torch.Tensor) -> None:
|
||||
"""Build the quantized representation of weight_fp32 and attach it to the
|
||||
module (register_buffer for tensors, plain attributes for scalars). Called
|
||||
while the module is still an nn.Linear, before the weight param is removed."""
|
||||
raise NotImplementedError
|
||||
|
||||
def dequantize(self, module: "OstrisLinear") -> torch.Tensor:
|
||||
"""Reconstruct the full weight in the original basis, in float32."""
|
||||
raise NotImplementedError
|
||||
|
||||
def requantize_(self, module: "OstrisLinear", fp_weight: torch.Tensor) -> None:
|
||||
"""Re-quantize in place from a full precision weight in the original basis
|
||||
(used by the continuous merge/reset method)."""
|
||||
raise NotImplementedError
|
||||
|
||||
def forward(self, module: "OstrisLinear", x: torch.Tensor) -> torch.Tensor:
|
||||
# default: dequantize per forward and run a plain linear. backends can
|
||||
# override with a cheaper formulation. the weight is frozen, so build it
|
||||
# outside autograd; gradients still flow to x through the matmul
|
||||
with torch.no_grad():
|
||||
w = self.dequantize(module).to(x.dtype)
|
||||
return F.linear(x, w, module.bias)
|
||||
|
||||
|
||||
class OstrisLinear(torch.nn.Linear):
|
||||
"""A linear layer whose weight is quantized by an OstrisQuantizer backend.
|
||||
|
||||
Never instantiate directly: created in place by convert_linear_to_ostris. The
|
||||
weight parameter is removed; the quantized representation lives in backend-owned
|
||||
buffers, plus:
|
||||
ostris_quantizer the backend instance
|
||||
ostris_orig_dtype dtype of the original weight (used for dequantized views)
|
||||
"""
|
||||
|
||||
is_ostris_quantized = True
|
||||
|
||||
@torch.no_grad()
|
||||
def dequantize_weight(self) -> torch.Tensor:
|
||||
"""Reconstruct the weight in the original basis and dtype."""
|
||||
return self.ostris_quantizer.dequantize(self).to(self.ostris_orig_dtype)
|
||||
|
||||
@property
|
||||
def weight(self):
|
||||
# materializes the full dequantized weight. kept for code that inspects the
|
||||
# weight (shape/dtype/device) and for the network merge paths, which detect
|
||||
# the marker via toolkit.util.quantize.is_quantized_tensor
|
||||
w = self.dequantize_weight()
|
||||
w._is_ostris_weight = True
|
||||
return w
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.ostris_quantizer.forward(self, x)
|
||||
|
||||
@torch.no_grad()
|
||||
def requantize_(self, fp_weight: torch.Tensor) -> None:
|
||||
self.ostris_quantizer.requantize_(self, fp_weight)
|
||||
|
||||
def _save_to_state_dict(self, destination, prefix, keep_vars):
|
||||
# emit a plain full precision weight so full-model saves need no special casing
|
||||
destination[prefix + "weight"] = self.dequantize_weight()
|
||||
if self.bias is not None:
|
||||
destination[prefix + "bias"] = self.bias if keep_vars else self.bias.detach()
|
||||
|
||||
|
||||
def get_ostris_quantizer(qtype: str) -> Optional[OstrisQuantizer]:
|
||||
"""Resolve a qtype string to a quantizer backend instance, or None if the qtype
|
||||
does not belong to a custom backend. Add new backends here."""
|
||||
from toolkit.util.orbit_quant import ORBIT_QTYPES, OrbitQuantizer
|
||||
from toolkit.util.orbit_vq_quant import ORBIT_VQ_QTYPES, OrbitVQQuantizer
|
||||
|
||||
if qtype in ORBIT_QTYPES:
|
||||
return OrbitQuantizer(ORBIT_QTYPES[qtype])
|
||||
if qtype in ORBIT_VQ_QTYPES:
|
||||
return OrbitVQQuantizer(**ORBIT_VQ_QTYPES[qtype])
|
||||
return None
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def convert_linear_to_ostris(module: torch.nn.Linear, quantizer: OstrisQuantizer) -> bool:
|
||||
"""Quantize an nn.Linear in place (class swap). Returns True if the module was
|
||||
converted (or already was), False if it is not a candidate."""
|
||||
if isinstance(module, OstrisLinear):
|
||||
return True
|
||||
weight = getattr(module, "weight", None)
|
||||
if not isinstance(weight, torch.nn.Parameter) or not weight.dtype.is_floating_point:
|
||||
return False
|
||||
if type(weight.data) is not torch.Tensor:
|
||||
# already holds a quantized tensor subclass (e.g. torchao)
|
||||
return False
|
||||
if not quantizer.can_quantize(module):
|
||||
return False
|
||||
quantizer.quantize_(module, weight.data.to(torch.float32))
|
||||
module.ostris_quantizer = quantizer
|
||||
module.ostris_orig_dtype = weight.dtype
|
||||
del module._parameters["weight"]
|
||||
if module.bias is not None:
|
||||
module.bias.requires_grad_(False)
|
||||
module.__class__ = OstrisLinear
|
||||
return True
|
||||
@@ -16,6 +16,12 @@ from safetensors.torch import load_file
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
from toolkit.print import print_acc
|
||||
from toolkit.util.ostris_quant import (
|
||||
OstrisLinear,
|
||||
OstrisQuantizer,
|
||||
convert_linear_to_ostris,
|
||||
get_ostris_quantizer,
|
||||
)
|
||||
import os
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -31,6 +37,7 @@ Q_MODULES = [
|
||||
"QLayerNorm",
|
||||
"QConvTranspose2d",
|
||||
"QEmbeddingBag",
|
||||
"OstrisLinear",
|
||||
]
|
||||
|
||||
torchao_qtypes = {
|
||||
@@ -53,10 +60,20 @@ class aotype:
|
||||
self.config = torchao_qtypes[name]
|
||||
|
||||
|
||||
class ostristype:
|
||||
# custom quantization backend (see toolkit/util/ostris_quant.py), e.g. orbit2/3/4
|
||||
def __init__(self, name: str, quantizer: OstrisQuantizer):
|
||||
self.name = name
|
||||
self.quantizer = quantizer
|
||||
|
||||
|
||||
def get_qtype(qtype: Union[str, qtype]) -> qtype:
|
||||
if qtype in torchao_qtypes:
|
||||
return aotype(qtype)
|
||||
if isinstance(qtype, str):
|
||||
ostris_quantizer = get_ostris_quantizer(qtype)
|
||||
if ostris_quantizer is not None:
|
||||
return ostristype(qtype, ostris_quantizer)
|
||||
return qtypes[qtype]
|
||||
else:
|
||||
return qtype
|
||||
@@ -65,6 +82,10 @@ def get_qtype(qtype: Union[str, qtype]) -> qtype:
|
||||
def is_quantized_tensor(t) -> bool:
|
||||
# torchao stores quantized weights as tensor subclasses (e.g. AffineQuantizedTensor) under torchao.*
|
||||
# that still report as nn.Parameter and expose .dequantize(). (quanto is handled separately.)
|
||||
# OstrisLinear.weight returns an already-dequantized tensor tagged with _is_ostris_weight
|
||||
# (its .dequantize() is a no-op) so the merge paths route through requantize_module_weight.
|
||||
if getattr(t, '_is_ostris_weight', False):
|
||||
return True
|
||||
return 'torchao' in type(t).__module__ and hasattr(t, 'dequantize')
|
||||
|
||||
|
||||
@@ -73,20 +94,33 @@ def dequantize_if_quantized(t):
|
||||
|
||||
|
||||
def get_torchao_config(qtype):
|
||||
# returns the torchao quantization config for a given qtype string, or None if it isn't torchao
|
||||
# returns the requantization config for a given qtype string (a torchao config, or the
|
||||
# ostristype for custom backends), or None if the qtype supports neither
|
||||
if qtype is None:
|
||||
return None
|
||||
try:
|
||||
q = get_qtype(qtype)
|
||||
except Exception:
|
||||
return None
|
||||
return q.config if isinstance(q, aotype) else None
|
||||
if isinstance(q, aotype):
|
||||
return q.config
|
||||
if isinstance(q, ostristype):
|
||||
return q
|
||||
return None
|
||||
|
||||
|
||||
def requantize_module_weight(module, fp_weight, orig_dtype, config) -> None:
|
||||
"""Write a full precision weight back into module.weight, re-quantizing in place if a torchao
|
||||
config is provided so the module stays quantized (used by the continuous merge/reset method).
|
||||
If config is None the weight is left in full precision."""
|
||||
"""Write a full precision weight back into module.weight, re-quantizing in place if a
|
||||
requantization config is provided so the module stays quantized (used by the continuous
|
||||
merge/reset method). If config is None the weight is left in full precision."""
|
||||
if isinstance(module, OstrisLinear):
|
||||
# the module's backend reuses its existing quantization state; config is not needed
|
||||
module.requantize_(fp_weight)
|
||||
return
|
||||
if isinstance(config, ostristype):
|
||||
# custom backend config but the module was never converted (e.g. skipped at
|
||||
# quantize time); leave it in full precision
|
||||
config = None
|
||||
module.weight = torch.nn.Parameter(fp_weight.to(orig_dtype), requires_grad=False)
|
||||
if config is not None:
|
||||
torchao_quantize_(module, config)
|
||||
@@ -142,7 +176,10 @@ def quantize(
|
||||
if m.__class__.__name__ in Q_MODULES:
|
||||
continue
|
||||
else:
|
||||
if isinstance(weights, aotype):
|
||||
if isinstance(weights, ostristype):
|
||||
if isinstance(m, torch.nn.Linear):
|
||||
convert_linear_to_ostris(m, weights.quantizer)
|
||||
elif isinstance(weights, aotype):
|
||||
torchao_quantize_(m, weights.config)
|
||||
else:
|
||||
_quantize_submodule(
|
||||
|
||||
Reference in New Issue
Block a user