Move uintx quantization to ostris quant with bit identical matching. Now we are not bound to an older version of torch ao.
This commit is contained in:
@@ -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"]
|
||||
|
||||
@@ -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(),
|
||||
}
|
||||
|
||||
205
toolkit/util/uintx_quant.py
Normal file
205
toolkit/util/uintx_quant.py
Normal file
@@ -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)
|
||||
Reference in New Issue
Block a user