From 5804ef3b17309a391f226e308a90063cab52819f Mon Sep 17 00:00:00 2001 From: Andrea Fasoli Date: Tue, 8 Sep 2026 17:21:08 -0400 Subject: [PATCH 1/2] replace soon-deprecated quant functions Signed-off-by: Andrea Fasoli --- fms_mo/custom_ext_kernels/utils.py | 8 +++- fms_mo/modules/bmm.py | 8 +++- fms_mo/modules/linear.py | 16 ++++--- fms_mo/quant/quantizers.py | 40 +++++++++++------ fms_mo/quant_refactor/base_quant.py | 16 +++++++ fms_mo/quant_refactor/pactplussym_rc.py | 9 ++-- fms_mo/quant_refactor/per_channel_ste.py | 22 +++++++--- fms_mo/quant_refactor/per_tensor_ste.py | 11 ++++- fms_mo/quant_refactor/quantizers_new.py | 55 +++++++++++++++++------- fms_mo/quant_refactor/torch_quantizer.py | 53 ++++++++++++++++------- 10 files changed, 176 insertions(+), 62 deletions(-) diff --git a/fms_mo/custom_ext_kernels/utils.py b/fms_mo/custom_ext_kernels/utils.py index 386b8c14..74ed8b46 100644 --- a/fms_mo/custom_ext_kernels/utils.py +++ b/fms_mo/custom_ext_kernels/utils.py @@ -791,7 +791,13 @@ def q_iaddmm_dq_abstract(bias, m1, m2, scale_i, zp_i, scale_w): @kernel_impl("fms_mo::q_per_t_sym", "default") def q_per_t(x, s, zp): - return torch.quantize_per_tensor(x, s, zp, torch.qint8).int_repr() + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); this + # mirrors it (reciprocal multiply in fp32, round half-to-even, saturate to int8). + return ( + (torch.round(x * torch.as_tensor(s).to(x.dtype).reciprocal()) + zp) + .clamp(-128, 127) + .to(torch.int8) + ) @reg_fake("fms_mo::q_per_t_sym") def q_per_t_abstract(x, s, zp): diff --git a/fms_mo/modules/bmm.py b/fms_mo/modules/bmm.py index 114a2df7..afa7be8c 100644 --- a/fms_mo/modules/bmm.py +++ b/fms_mo/modules/bmm.py @@ -685,7 +685,13 @@ def qfunc_pt(self, x, scale, zp): Returns: Tensor: The quantized tensor. """ - return torch.quantize_per_tensor(x.float(), scale, zp, torch.qint8).int_repr() + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); this + # mirrors it (reciprocal multiply in fp32, round half-to-even, saturate to int8). + return ( + (torch.round(x.float() * torch.as_tensor(scale).float().reciprocal()) + zp) + .clamp(-128, 127) + .to(torch.int8) + ) def qfunc_raw(self, x, scale, zp): """ diff --git a/fms_mo/modules/linear.py b/fms_mo/modules/linear.py index 3a39bb30..4e596d32 100644 --- a/fms_mo/modules/linear.py +++ b/fms_mo/modules/linear.py @@ -1051,12 +1051,16 @@ def qa_pt_quant_func(self, x): Returns: Tensor: Quantized tensor with values in the range [-128, 127]. """ - return torch.quantize_per_tensor( - x.float(), - self.input_scale, - self.input_zp - 128 + self.useSymAct, - torch.qint8, - ).int_repr() + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); this + # mirrors it (reciprocal multiply in fp32, round half-to-even, saturate to int8). + return ( + ( + torch.round(x.float() * self.input_scale.float().reciprocal()) + + (self.input_zp - 128 + self.useSymAct) + ) + .clamp(-128, 127) + .to(torch.int8) + ) def qa_raw_qfunc(self, x): """ diff --git a/fms_mo/quant/quantizers.py b/fms_mo/quant/quantizers.py index f356eb6c..e5c04c32 100644 --- a/fms_mo/quant/quantizers.py +++ b/fms_mo/quant/quantizers.py @@ -42,6 +42,16 @@ logger = logging.getLogger(__name__) +# dtype that .int_repr() returned for each deprecated quantized dtype, used to reproduce +# torch.quantize_per_tensor/_per_channel without the deprecated ops +# (pytorch/pytorch#184982). +_INT_REPR_DTYPES = { + torch.qint32: torch.int32, + torch.qint8: torch.int8, + torch.quint8: torch.uint8, + torch.int32: torch.int32, +} + def get_activation_quantizer( qa_mode="PACT", @@ -555,12 +565,14 @@ def forward( clip_val.dtype ) # NOTE return will be a fp32 tensor; function only support float() else: + # NOTE torch.quantize_per_channel is deprecated (pytorch/pytorch#184982), use + # plain arithmetic instead. zero_point is always 0 here, so it drops out. + # torch.round matches the deprecated op's round-half-to-even tie-breaking. + scale_bcast = scale.reshape([-1] + [1] * (input_tensor.dim() - 1)) output = ( - torch.quantize_per_channel( - input_tensor, scale, zero_point, 0, torch.qint8 - ) - .int_repr() + torch.round(input_tensor / scale_bcast) .clamp(int_l, int_u) + .to(torch.int8) ) # NOTE return will be a torch.int8 tensor @@ -1563,13 +1575,15 @@ def forward( ) out = out.to(input_tensor_dtype) else: + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); + # quantize with plain arithmetic instead. The deprecated kernel multiplied + # by the reciprocal of the scale in fp32 and rounded half-to-even; mirror + # that exactly, since these results are compared against TorchQuantizer. # Clamp to [quant_min, quant_max] in case we are storing int4 into a uint8 tensor out = ( - torch.quantize_per_tensor( - input_tensor, scale.float(), zp, qint_dtype - ) - .int_repr() + (torch.round(input_tensor * scale.float().reciprocal()) + zp) .clamp(quant_min, quant_max) + .to(_INT_REPR_DTYPES[qint_dtype]) ) return out # NOTE remember scale and zp from asym_lin_q_params is different from @@ -2030,13 +2044,15 @@ def forward( ) out = out.to(input_tensor_dtype) else: + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); + # quantize with plain arithmetic instead. The deprecated kernel multiplied + # by the reciprocal of the scale in fp32 and rounded half-to-even; mirror + # that exactly, since these results are compared against TorchQuantizer. # Clamp to [quant_min, quant_max] in case we are storing quint4 into a uint8 tensor out = ( - torch.quantize_per_tensor( - input_tensor, scale.float(), zp, qint_dtype - ) - .int_repr() + (torch.round(input_tensor * scale.float().reciprocal()) + zp) .clamp(quant_min, quant_max) + .to(_INT_REPR_DTYPES[qint_dtype]) ) return out # do not cast back to input_tensor_dtype! diff --git a/fms_mo/quant_refactor/base_quant.py b/fms_mo/quant_refactor/base_quant.py index 0960ff32..1a0ac495 100644 --- a/fms_mo/quant_refactor/base_quant.py +++ b/fms_mo/quant_refactor/base_quant.py @@ -26,6 +26,22 @@ # Third Party import torch +# The deprecated torch.quantize_per_tensor/_per_channel ops (pytorch/pytorch#184982) +# saturated to the storage dtype's range and int_repr() returned these dtypes. Kept here +# so the plain-arithmetic replacements reproduce that behavior exactly. +_INT_REPR_DTYPES = { + torch.qint32: torch.int32, + torch.qint8: torch.int8, + torch.quint8: torch.uint8, + torch.int32: torch.int32, +} +_DTYPE_RANGES = { + torch.qint32: (-(2**31), 2**31 - 1), + torch.qint8: (-128, 127), + torch.quint8: (0, 255), + torch.int32: (-(2**31), 2**31 - 1), +} + @dataclass class Qscheme: diff --git a/fms_mo/quant_refactor/pactplussym_rc.py b/fms_mo/quant_refactor/pactplussym_rc.py index 145e11b5..bc2c37b2 100644 --- a/fms_mo/quant_refactor/pactplussym_rc.py +++ b/fms_mo/quant_refactor/pactplussym_rc.py @@ -312,10 +312,13 @@ def forward( quant_max=qint_h, ).to(input_tensor.dtype) else: - qint_dtype = torch.qint8 + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); + # mirror it with plain arithmetic (reciprocal multiply in fp32, round + # half-to-even, saturate to int8) instead. output = ( - torch.quantize_per_tensor(input_tensor, scale, zero_point, qint_dtype) - .int_repr() + (torch.round(input_tensor * scale.reciprocal()) + zero_point) + .clamp(-128, 127) + .to(torch.int8) .clamp(qint_l, qint_h) ) return output diff --git a/fms_mo/quant_refactor/per_channel_ste.py b/fms_mo/quant_refactor/per_channel_ste.py index 6b56a5ff..cd26c30c 100644 --- a/fms_mo/quant_refactor/per_channel_ste.py +++ b/fms_mo/quant_refactor/per_channel_ste.py @@ -23,6 +23,7 @@ import torch # Local +from fms_mo.quant_refactor.base_quant import _DTYPE_RANGES, _INT_REPR_DTYPES from fms_mo.quant_refactor.linear_utils import ( asymmetric_linear_quantization_params, linear_quantization, @@ -361,15 +362,22 @@ def linear_quantization( quant_max=qint_h, ).to(input_tensor.dtype) else: + # NOTE torch.quantize_per_channel is deprecated (pytorch/pytorch#184982); + # mirror it with plain arithmetic (reciprocal multiply in fp32, round + # half-to-even, saturate to the storage dtype) instead. + bcast = [1] * input_tensor.dim() + bcast[axis] = -1 + dtype_l, dtype_h = _DTYPE_RANGES[qint_dtype] output = ( - torch.quantize_per_channel( - input_tensor.float(), - scale.float(), - zero_point, - axis=axis, - dtype=qint_dtype, + ( + torch.round( + input_tensor.float() * scale.float().reciprocal().reshape(bcast) + ) + + zero_point.reshape(bcast) ) - .int_repr() + .to(torch.float64) + .clamp(dtype_l, dtype_h) + .to(_INT_REPR_DTYPES[qint_dtype]) .clamp(qint_l, qint_h) ) return output diff --git a/fms_mo/quant_refactor/per_tensor_ste.py b/fms_mo/quant_refactor/per_tensor_ste.py index 32bc9cdf..9261cf5c 100644 --- a/fms_mo/quant_refactor/per_tensor_ste.py +++ b/fms_mo/quant_refactor/per_tensor_ste.py @@ -26,6 +26,7 @@ import torch # Local +from fms_mo.quant_refactor.base_quant import _DTYPE_RANGES, _INT_REPR_DTYPES from fms_mo.quant_refactor.linear_utils import ( asymmetric_linear_quantization_params, linear_quantization, @@ -335,9 +336,15 @@ def linear_quantization( quant_max=qint_h, ).to(input_tensor.dtype) else: + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); + # mirror it with plain arithmetic (reciprocal multiply in fp32, round + # half-to-even, saturate to the storage dtype) instead. + dtype_l, dtype_h = _DTYPE_RANGES[qint_dtype] output = ( - torch.quantize_per_tensor(input_tensor, scale, zero_point, qint_dtype) - .int_repr() + (torch.round(input_tensor * scale.reciprocal()) + zero_point) + .to(torch.float64) + .clamp(dtype_l, dtype_h) + .to(_INT_REPR_DTYPES[qint_dtype]) .clamp(qint_l, qint_h) ) return output diff --git a/fms_mo/quant_refactor/quantizers_new.py b/fms_mo/quant_refactor/quantizers_new.py index d54747de..1a2d2299 100644 --- a/fms_mo/quant_refactor/quantizers_new.py +++ b/fms_mo/quant_refactor/quantizers_new.py @@ -35,6 +35,9 @@ import torch.nn as nn import torch.nn.functional as F +# Local +from fms_mo.quant_refactor.base_quant import _DTYPE_RANGES, _INT_REPR_DTYPES + logger = logging.getLogger(__name__) @@ -557,9 +560,17 @@ def forward(ctx, input, num_bits, dequantize, inplace, objSAWB_clip_val): clip_val.dtype ) # NOTE return will be a fp32 tensor; function only support float() else: + # NOTE torch.quantize_per_channel is deprecated (pytorch/pytorch#184982); + # mirror it with plain arithmetic (reciprocal multiply in fp32, round + # half-to-even, saturate to int8) instead. + bcast = [-1] + [1] * (input.dim() - 1) output = ( - torch.quantize_per_channel(input, scale, zero_point, 0, torch.qint8) - .int_repr() + ( + torch.round(input * scale.reciprocal().reshape(bcast)) + + zero_point.reshape(bcast) + ) + .clamp(-128, 127) + .to(torch.int8) .clamp(int_l, int_u) ) # NOTE return will be a torch.int8 tensor @@ -1561,9 +1572,15 @@ def forward( return out.to(input_dtype) else: # Clamp to [quant_min, quant_max] in case we are storing int4 into a uint8 tensor + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); + # mirror it with plain arithmetic (reciprocal multiply in fp32, round + # half-to-even, saturate to the storage dtype) instead. + dtype_l, dtype_h = _DTYPE_RANGES[qint_dtype] out = ( - torch.quantize_per_tensor(input, scale.float(), zp, qint_dtype) - .int_repr() + (torch.round(input * scale.float().reciprocal()) + zp) + .to(torch.float64) + .clamp(dtype_l, dtype_h) + .to(_INT_REPR_DTYPES[qint_dtype]) .clamp(quant_min, quant_max) ) return out # do not cast back to input_dtype! @@ -1979,11 +1996,15 @@ def forward( return out.to(input_dtype) else: # Clamp to [quant_min, quant_max] in case we are storing quint4 into a uint8 tensor + # NOTE torch.quantize_per_tensor is deprecated (pytorch/pytorch#184982); + # mirror it with plain arithmetic (reciprocal multiply in fp32, round + # half-to-even, saturate to the storage dtype) instead. + dtype_l, dtype_h = _DTYPE_RANGES[qint_dtype] out = ( - torch.quantize_per_tensor( - input, scale.float(), zero_point, qint_dtype - ) - .int_repr() + (torch.round(input * scale.float().reciprocal()) + zero_point) + .to(torch.float64) + .clamp(dtype_l, dtype_h) + .to(_INT_REPR_DTYPES[qint_dtype]) .clamp(quant_min, quant_max) ) return out # do not cast back to input_dtype! @@ -3122,15 +3143,19 @@ def forward( quant_max=int_u, ).to(input.dtype) else: + # NOTE torch.quantize_per_channel is deprecated (pytorch/pytorch#184982); + # mirror it with plain arithmetic (reciprocal multiply in fp32, round + # half-to-even, saturate to int8) instead. + bcast = [-1] + [1] * (input.dim() - 1) output = ( - torch.quantize_per_channel( - input.float(), - scale.float(), - zero_point.float(), - axis=0, - dtype=torch.qint8, + ( + torch.round( + input.float() * scale.float().reciprocal().reshape(bcast) + ) + + zero_point.float().reshape(bcast) ) - .int_repr() + .clamp(-128, 127) + .to(torch.int8) .clamp(int_l, int_u) ) diff --git a/fms_mo/quant_refactor/torch_quantizer.py b/fms_mo/quant_refactor/torch_quantizer.py index 7dd4da71..de2d7f3d 100644 --- a/fms_mo/quant_refactor/torch_quantizer.py +++ b/fms_mo/quant_refactor/torch_quantizer.py @@ -26,7 +26,11 @@ import torch # Local -from fms_mo.quant_refactor.base_quant import Qscheme +from fms_mo.quant_refactor.base_quant import ( + _DTYPE_RANGES, + _INT_REPR_DTYPES, + Qscheme, +) from fms_mo.quant_refactor.sawb_utils import sawb_params, sawb_params_code logger = logging.getLogger(__name__) @@ -277,6 +281,20 @@ def get_torch_dtype(self): (self.num_bits_int, signed) ) # NOTE .item() won't work for perCh + def _broadcast_qparams(self, ndim: int): + """ + Reshape perCh scale/zero_point so they broadcast along qscheme.axis. + + Args: + ndim (int): Number of dimensions of the tensor being quantized. + + Returns: + [torch.Tensor, torch.Tensor]: Broadcastable scale and zero_point. + """ + shape = [1] * ndim + shape[self.qscheme.axis] = -1 + return self.scale.reshape(shape), self.zero_point.reshape(shape) + def forward(self, tensor: torch.Tensor): """ TorchQuantizer forward() function w/ PT kernels. @@ -314,27 +332,32 @@ def forward(self, tensor: torch.Tensor): else: dtype = self.get_torch_dtype() if dtype: + # NOTE torch.quantize_per_tensor/_per_channel are deprecated + # (pytorch/pytorch#184982), so quantize with plain arithmetic instead. + # This class is the reference result for the quantizer tests, so the + # arithmetic mirrors the deprecated kernels bit-for-bit: they multiply + # by the reciprocal of the scale in fp32 (NOT divide -- that is slightly + # more accurate and would shift the reference), round half-to-even, and + # saturate to the storage dtype range before the clamp below. if self.qscheme.q_unit == "perCh": - output = torch.quantize_per_channel( - tensor, - self.scale, - self.zero_point, - self.qscheme.axis, - dtype, - ) + scale, zero_point = self._broadcast_qparams(tensor.dim()) elif self.qscheme.q_unit == "perGrp": raise RuntimeError( "TorchQuantizer forward not implemented for perGrp" ) else: # Per Tensor - output = torch.quantize_per_tensor( - tensor, - self.scale, - self.zero_point, - dtype, - ) + scale, zero_point = self.scale, self.zero_point + dtype_min, dtype_max = _DTYPE_RANGES[dtype] + # NOTE clamp in float64: qint32's bounds are not representable in + # float32, so clamping there would overflow on cast. + output = ( + (torch.round(tensor * scale.reciprocal()) + zero_point) + .to(torch.float64) + .clamp(dtype_min, dtype_max) + .to(_INT_REPR_DTYPES[dtype]) + ) # Clamp required if storing int4 into int8 tensor (no PT support for int4) - output = output.int_repr().clamp(self.quant_min, self.quant_max) + output = output.clamp(self.quant_min, self.quant_max) else: raise RuntimeError( f"num_bits {self.num_bits} and sign {(self.zero_point==0).item()}" From 32391fb6668d042c9d32862ba5ba899e355eec5a Mon Sep 17 00:00:00 2001 From: Andrea Fasoli Date: Tue, 8 Sep 2026 17:43:59 -0400 Subject: [PATCH 2/2] ruff fixes Signed-off-by: Andrea Fasoli --- fms_mo/quant_refactor/torch_quantizer.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/fms_mo/quant_refactor/torch_quantizer.py b/fms_mo/quant_refactor/torch_quantizer.py index de2d7f3d..f5d90e16 100644 --- a/fms_mo/quant_refactor/torch_quantizer.py +++ b/fms_mo/quant_refactor/torch_quantizer.py @@ -26,11 +26,7 @@ import torch # Local -from fms_mo.quant_refactor.base_quant import ( - _DTYPE_RANGES, - _INT_REPR_DTYPES, - Qscheme, -) +from fms_mo.quant_refactor.base_quant import _DTYPE_RANGES, _INT_REPR_DTYPES, Qscheme from fms_mo.quant_refactor.sawb_utils import sawb_params, sawb_params_code logger = logging.getLogger(__name__)