Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion fms_mo/custom_ext_kernels/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
8 changes: 7 additions & 1 deletion fms_mo/modules/bmm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
"""
Expand Down
16 changes: 10 additions & 6 deletions fms_mo/modules/linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
"""
Expand Down
40 changes: 28 additions & 12 deletions fms_mo/quant/quantizers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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!

Expand Down
16 changes: 16 additions & 0 deletions fms_mo/quant_refactor/base_quant.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
9 changes: 6 additions & 3 deletions fms_mo/quant_refactor/pactplussym_rc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
22 changes: 15 additions & 7 deletions fms_mo/quant_refactor/per_channel_ste.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
11 changes: 9 additions & 2 deletions fms_mo/quant_refactor/per_tensor_ste.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
55 changes: 40 additions & 15 deletions fms_mo/quant_refactor/quantizers_new.py
Original file line number Diff line number Diff line change
Expand Up @@ -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__)


Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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!
Expand Down Expand Up @@ -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!
Expand Down Expand Up @@ -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)
)

Expand Down
Loading
Loading