From b677cdb02666320f1b03c747f5037a41e5a7515e Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Mon, 27 Jul 2026 12:59:42 -0600 Subject: [PATCH] Move uintx quantization to ostris quant with bit identical matching. Now we are not bound to an older version of torch ao. --- toolkit/util/ostris_quant.py | 15 ++- toolkit/util/quantize.py | 11 +- toolkit/util/uintx_quant.py | 205 +++++++++++++++++++++++++++++++++++ 3 files changed, 221 insertions(+), 10 deletions(-) create mode 100644 toolkit/util/uintx_quant.py diff --git a/toolkit/util/ostris_quant.py b/toolkit/util/ostris_quant.py index ae03f54..a647295 100644 --- a/toolkit/util/ostris_quant.py +++ b/toolkit/util/ostris_quant.py @@ -32,6 +32,11 @@ class OstrisQuantizer: # get_ostris_quantizer); quantized saves need it to restore the backend qtype: Optional[str] = None + # backends that quantize in the weight's own dtype can set this False to + # receive the raw weight tensor in quantize_ instead of a float32 copy, + # avoiding a 2x-weight-size allocation during model quantization + wants_fp32_weight: bool = True + def can_quantize(self, module: torch.nn.Linear) -> bool: """Whether this backend can quantize the given linear (e.g. shape constraints).""" return True @@ -173,9 +178,10 @@ class OstrisLazyWeight(torch.Tensor): 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.convrot_quant import CONVROT_QTYPES, get_convrot_quantizer from toolkit.util.orbit_quant import ORBIT_QTYPES, OrbitQuantizer from toolkit.util.orbit_vq_quant import ORBIT_VQ_QTYPES, OrbitVQQuantizer - from toolkit.util.convrot_quant import CONVROT_QTYPES, get_convrot_quantizer + from toolkit.util.uintx_quant import UINTX_QTYPES, UIntXQuantizer quantizer = None if qtype in ORBIT_QTYPES: @@ -184,6 +190,8 @@ def get_ostris_quantizer(qtype: str) -> Optional[OstrisQuantizer]: quantizer = OrbitVQQuantizer(**ORBIT_VQ_QTYPES[qtype]) elif qtype in CONVROT_QTYPES: quantizer = get_convrot_quantizer(qtype) + elif qtype in UINTX_QTYPES: + quantizer = UIntXQuantizer(UINTX_QTYPES[qtype]) if quantizer is not None: # quantized saves read this back to restore the backend on load quantizer.qtype = qtype @@ -327,7 +335,10 @@ def convert_linear_to_ostris( return False if not quantizer.can_quantize(module): return False - quantizer.quantize_(module, weight.data.to(torch.float32)) + if quantizer.wants_fp32_weight: + quantizer.quantize_(module, weight.data.to(torch.float32)) + else: + quantizer.quantize_(module, weight.data) module.ostris_quantizer = quantizer module.ostris_orig_dtype = weight.dtype del module._parameters["weight"] diff --git a/toolkit/util/quantize.py b/toolkit/util/quantize.py index 2f7da71..9307f3d 100644 --- a/toolkit/util/quantize.py +++ b/toolkit/util/quantize.py @@ -7,7 +7,6 @@ from optimum.quanto.tensor import Optimizer, qtype, qtypes from torchao.quantization.quant_api import ( quantize_ as torchao_quantize_, Float8WeightOnlyConfig, - UIntXWeightOnlyConfig, Int8WeightOnlyConfig ) from optimum.quanto import freeze @@ -42,13 +41,9 @@ Q_MODULES = [ torchao_qtypes = { # "int4": Int4WeightOnlyConfig(), - "uint2": UIntXWeightOnlyConfig(torch.uint2), - "uint3": UIntXWeightOnlyConfig(torch.uint3), - "uint4": UIntXWeightOnlyConfig(torch.uint4), - "uint5": UIntXWeightOnlyConfig(torch.uint5), - "uint6": UIntXWeightOnlyConfig(torch.uint6), - "uint7": UIntXWeightOnlyConfig(torch.uint7), - "uint8": UIntXWeightOnlyConfig(torch.uint8), + # uint2..uint8 are handled by the UIntXQuantizer ostris backend + # (toolkit/util/uintx_quant.py), a bit-exact reproduction of torchao 0.10.0's + # UIntXWeightOnlyConfig, so ARAs stay byte-identical after torchao upgrades "int8": Int8WeightOnlyConfig(), "float8": Float8WeightOnlyConfig(), } diff --git a/toolkit/util/uintx_quant.py b/toolkit/util/uintx_quant.py new file mode 100644 index 0000000..4c123a3 --- /dev/null +++ b/toolkit/util/uintx_quant.py @@ -0,0 +1,205 @@ +""" +Bit-exact reimplementation of torchao 0.10.0's UIntXWeightOnlyConfig weight +quantization as an OstrisQuantizer backend, so the uint2..uint7 qtypes keep +producing byte-identical weights after torchao drops uintx support. Existing +accuracy recovery adapters were trained against these exact quantized bases, +so every op below mirrors torchao's sequence (same order, same dtypes): + + choose_qparams_affine (ASYMMETRIC, block_size (1, 64), preserve_zero=True, + eps=float32 eps, INT zero-point domain, scale in the weight dtype): + min/max per 64-wide group along in_features, extended to include 0 + scale = (max_val_pos - min_val_neg) / (qmax - qmin), clamped to eps + zero_point = clamp(qmin - round(min_val_neg / scale), qmin, qmax) as int + quantize_affine: + q = clamp(round(w * (1.0 / scale)) + zero_point, qmin, qmax) + dequantize_affine: + w = ((q as int) - zero_point) cast to the weight dtype, times scale + +The arithmetic is done in the original weight dtype (usually bfloat16) because +that is what torchao did; doing it in float32 would round differently. + +uint8 resolves to a backend too, but can_quantize always refuses it: torchao +0.10.0's uint8 path raised inside UintxTensor.from_uint8 (packing only supports +1..7 bits) before the module was touched, so "uint8" layers have always been +silently left unquantized. Refusing keeps that exact behavior — flip the check +in can_quantize if real uint8 quantization is ever wanted for new models. + +Buffers on the module (registered by quantize_): + uintx_packed quantized codes packed into power-of-2 bit shards (like + torchao's UintxTensor: e.g. uint3 = a 2-bit + a 1-bit + shard), concatenated into one flat uint8 buffer. Shards + unpack with a couple of elementwise shift/mask kernels, + which is much cheaper per forward than a generic bitstream. + uintx_scale per-group scale, stored as a uint8 byte view of the weight + dtype so module.to(dtype=...) can't cast it + uintx_zero_point per-group zero point, uint8 (values are in [0, qmax]) +""" + +import torch + +from toolkit.util.ostris_quant import OstrisLinear, OstrisQuantizer + +UINTX_QTYPES = {f"uint{bits}": bits for bits in range(2, 9)} + +_EPS = torch.finfo(torch.float32).eps + + +def _pack_shard(vals: torch.Tensor, k: int) -> torch.Tensor: + """Pack flat uint8 values (< 2**k) into bytes, 8 // k values per byte. + Values are laid out in 8//k contiguous chunks (chunk j holds bits + [j*k, j*k+k) of every byte) so pack and unpack touch memory coalesced.""" + vpb = 8 // k + if vpb == 1: + return vals.clone() + pad = (-vals.numel()) % vpb + if pad: + vals = torch.cat([vals, vals.new_zeros(pad)]) + chunks = vals.view(vpb, -1) + out = chunks[0].clone() + for j in range(1, vpb): + out |= chunks[j] << (j * k) + return out + + +def _unpack_shard(packed: torch.Tensor, k: int, numel: int) -> torch.Tensor: + vpb = 8 // k + if vpb == 1: + return packed[:numel] + out = torch.empty(vpb, packed.numel(), dtype=torch.uint8, device=packed.device) + for j in range(vpb): + torch.bitwise_right_shift(packed, j * k, out=out[j]) + out.bitwise_and_((1 << k) - 1) + return out.view(-1)[:numel] + + +def pack_uintx(codes: torch.Tensor, nbits: int) -> torch.Tensor: + """Pack integer codes (values < 2**nbits) into concatenated power-of-2 bit + shards; bits [offset, offset+k) of each code land in the k-bit shard.""" + flat = codes.flatten().to(torch.uint8) + shards = [] + offset = 0 + for k in (8, 4, 2, 1): + if nbits & k: + if k == nbits: # single shard, values are already < 2**k + shards.append(_pack_shard(flat, k)) + else: + shards.append(_pack_shard((flat >> offset) & ((1 << k) - 1), k)) + offset += k + return torch.cat(shards) if len(shards) > 1 else shards[0] + + +def unpack_uintx(packed: torch.Tensor, nbits: int, numel: int) -> torch.Tensor: + """Inverse of pack_uintx. Returns a flat uint8 tensor of length numel.""" + out = None + offset = 0 + pos = 0 + for k in (8, 4, 2, 1): + if nbits & k: + vpb = 8 // k + nbytes = -(-numel // vpb) + vals = _unpack_shard(packed[pos : pos + nbytes], k, numel) + if out is None: + # a multi-shard first value is always a fresh tensor from + # _unpack_shard, safe to mutate; a lone 8-bit shard is a view + # of the buffer but is never combined, so read-only is fine + out = vals + else: + out |= vals.bitwise_left_shift_(offset) + offset += k + pos += nbytes + return out + + +class UIntXQuantizer(OstrisQuantizer): + # quantization runs in the weight's own dtype (that is what torchao did); + # skipping the float32 copy keeps peak memory at torchao levels + wants_fp32_weight = False + + def __init__(self, nbits: int, group_size: int = 64): + self.nbits = nbits + self.group_size = group_size + self.qmin = 0 + self.qmax = (1 << nbits) - 1 + + def can_quantize(self, module: torch.nn.Linear) -> bool: + if self.nbits == 8: + # see module docstring: uint8 has always been a silent no-op + return False + weight = getattr(module, "weight", None) + if weight is None or weight.dim() != 2: + return False + # torchao asserted on non-divisible in_features, which the quantize loop + # caught, leaving the layer unquantized; refuse for the same end state + return weight.shape[1] % self.group_size == 0 + + @torch.no_grad() + def quantize_(self, module: torch.nn.Linear, weight: torch.Tensor) -> None: + # weight arrives in its original dtype (wants_fp32_weight is False) + self._quantize_impl(module, weight) + + @torch.no_grad() + def requantize_(self, module: "OstrisLinear", fp_weight: torch.Tensor) -> None: + # the torchao path cast the merged weight to the original dtype and + # re-quantized in that dtype + self._quantize_impl(module, fp_weight.to(module.ostris_orig_dtype)) + + def _quantize_impl(self, module: torch.nn.Module, w: torch.Tensor) -> None: + out_f, in_f = w.shape + groups = in_f // self.group_size + wv = w.contiguous().view(out_f, groups, self.group_size) + + min_val = torch.amin(wv, dim=2) + max_val = torch.amax(wv, dim=2) + # preserve_zero: the qparams must be able to represent 0.0 exactly + min_val_neg = torch.min(min_val, torch.zeros_like(min_val)) + max_val_pos = torch.max(max_val, torch.zeros_like(max_val)) + + scale = (max_val_pos - min_val_neg) / float(self.qmax - self.qmin) + scale = torch.clamp(scale, min=_EPS) + zero_point = self.qmin - torch.round(min_val_neg / scale) + zero_point = torch.clamp(zero_point, self.qmin, self.qmax).to(torch.int32) + + q = torch.clamp( + torch.round(wv * (1.0 / scale.view(out_f, groups, 1))) + + zero_point.view(out_f, groups, 1), + self.qmin, + self.qmax, + ) + q = q.view(out_f, in_f).to(torch.uint8) + + module.register_buffer( + "uintx_packed", pack_uintx(q, self.nbits), persistent=False + ) + module.register_buffer( + "uintx_scale", scale.contiguous().view(torch.uint8), persistent=False + ) + module.register_buffer( + "uintx_zero_point", zero_point.to(torch.uint8), persistent=False + ) + + def _dequantize_native(self, module: "OstrisLinear") -> torch.Tensor: + """Reconstruct the weight in its original dtype. torchao subtracted the + zero point in int32 then cast; doing it directly in the weight dtype is + bit-identical (codes and zero points are integers <= 255, and every + integer in [-255, 255] is exact in bf16/fp16/fp32) with half the + intermediate memory traffic.""" + out_f, in_f = module.out_features, module.in_features + groups = in_f // self.group_size + scale = module.uintx_scale.view(module.ostris_orig_dtype) + q = unpack_uintx(module.uintx_packed, self.nbits, out_f * in_f) + dq = q.view(out_f, groups, self.group_size).to(scale.dtype) + dq -= module.uintx_zero_point.to(scale.dtype).view(out_f, groups, 1) + dq *= scale.view(out_f, groups, 1) + return dq.view(out_f, in_f) + + def dequantize(self, module: "OstrisLinear") -> torch.Tensor: + return self._dequantize_native(module).to(torch.float32) + + def forward(self, module: "OstrisLinear", x: torch.Tensor) -> torch.Tensor: + # skip the float32 round-trip of the base implementation; the weight is + # frozen, so build it outside autograd + with torch.no_grad(): + w = self._dequantize_native(module) + if w.dtype != x.dtype: + w = w.to(x.dtype) + return torch.nn.functional.linear(x, w, module.bias)