Summary
torchao 0.18.0+xpu rewrote the quantized tensor subclasses (Int8Tensor / Float8Tensor) to use __torch_dispatch__ tables, but several operators are missing from the dispatch registrations of exactly these two classes. This breaks quantize_() whenever the flow calls aten.view on an Int8Tensor or aten.abs/arithmetic ops on a Float8Tensor — e.g. re-quantizing an already-quantized module goes through _choose_qparams_affine which calls input.view(...).
Worse, the failure is swallowed by per-module try/except in callers: the module is left unquantized and only a Failed to quantize ... message is printed, so users believe quantization is on when it isn't.
Environment
- OS: Windows 11
- Python: 3.12.9
- torch:
2.13.0+xpu (Intel Arc A770)
- torchao:
0.18.0+xpu (from https://download.pytorch.org/whl/xpu)
Reproduction
import torch
from torchao.quantization.quant_api import quantize_ as torchao_quantize_
from torchao.quantization.quant_api import Int8WeightOnlyConfig, Float8WeightOnlyConfig
model = torch.nn.Sequential(torch.nn.Linear(64, 128), torch.nn.Linear(128, 16)).to("xpu")
# First pass: OK (weights become Int8Tensor)
torchao_quantize_(model, Int8WeightOnlyConfig())
# Second pass on an already-quantized module (what ai-toolkit's quantize() does
# when it quantizes the container, then re-visits each child Linear):
torchao_quantize_(model[0], Int8WeightOnlyConfig())
# NotImplementedError: Int8Tensor dispatch: attempting to run unimplemented
# operator/function: func=<OpOverload(op='aten.view', overload='default')>, ...
# Direct repro:
model[0].weight.view(64, 128) # same NotImplementedError
Float8Tensor equivalent:
model = torch.nn.Sequential(torch.nn.Linear(64, 128), torch.nn.Linear(128, 16)).to("xpu")
torchao_quantize_(model, Float8WeightOnlyConfig())
model[0].weight.abs() # NotImplementedError: Float8Tensor dispatch: aten.abs unimplemented
torchao_quantize_(model[0], Float8WeightOnlyConfig())
# after registering abs, next missing: aten.div (see below)
Full stack for int8:
torchao_quantize_ -> ... -> _choose_qparams_affine -> input.view(shape_for_reduction)
-> torchao.utils._dispatch__torch_dispatch__ -> NotImplementedError
Root cause analysis (dispatch tables in 0.18.0)
torchao/quantization/quantize_/workflows/int8/int8_tensor.py registers only:
aten.linear, aten.slice, aten.index, aten.embedding, aten.is_pinned,
aten._pin_memory, aten.select — no aten.view.
torchao/quantization/quantize_/workflows/float8/float8_tensor.py registers
aten.view (line ~863) but is missing aten.abs, and the requantize path
then also hits missing aten.div (and likely mul/add/sub/round/clamp/neg).
Sibling classes (Int4Tensor, NF4Tensor, NVFP4Tensor, MXTensor) all register
aten.view, which suggests this is an omission during the 0.18 rewrite rather than
an intentional design.
Because the dispatch tables are pure-Python per-class tables, this is not XPU-specific:
the same missing registrations should fail on CUDA as well (we only verified on XPU here).
Expected behavior
aten.view on Int8Tensor and aten.abs/basic arithmetic on Float8Tensor should
either be implemented (shape/view on the quantized data; abs/arith via dequantize
then re-quantize), or quantize_() should explicitly no-op on already-quantized
modules instead of failing per-module.
Workaround (verified working)
Runtime registration into the class dispatch tables works:
from torchao.quantization import Int8Tensor, Float8Tensor
@Int8Tensor.implements(torch.ops.aten.view.default)
def _view(func, types, args, kwargs):
t = args[0]
size = args[1] if len(args) > 1 else (kwargs or {}).get("size")
return t.dequantize().view(size)
@Float8Tensor.implements(torch.ops.aten.abs.default)
def _abs(func, types, args, kwargs):
return torch.abs(args[0].dequantize())
# plus div/mul/add/sub/round/clamp/neg with dequantize-based impls
With these registrations, both int8 and float8 quantize_() + forward pass on XPU.
Summary
torchao
0.18.0+xpurewrote the quantized tensor subclasses (Int8Tensor/Float8Tensor) to use__torch_dispatch__tables, but several operators are missing from the dispatch registrations of exactly these two classes. This breaksquantize_()whenever the flow callsaten.viewon anInt8Tensororaten.abs/arithmetic ops on aFloat8Tensor— e.g. re-quantizing an already-quantized module goes through_choose_qparams_affinewhich callsinput.view(...).Worse, the failure is swallowed by per-module try/except in callers: the module is left unquantized and only a
Failed to quantize ...message is printed, so users believe quantization is on when it isn't.Environment
2.13.0+xpu(Intel Arc A770)0.18.0+xpu(fromhttps://download.pytorch.org/whl/xpu)Reproduction
Float8Tensor equivalent:
Full stack for int8:
Root cause analysis (dispatch tables in 0.18.0)
torchao/quantization/quantize_/workflows/int8/int8_tensor.pyregisters only:aten.linear,aten.slice,aten.index,aten.embedding,aten.is_pinned,aten._pin_memory,aten.select— noaten.view.torchao/quantization/quantize_/workflows/float8/float8_tensor.pyregistersaten.view(line ~863) but is missingaten.abs, and the requantize paththen also hits missing
aten.div(and likelymul/add/sub/round/clamp/neg).Sibling classes (
Int4Tensor,NF4Tensor,NVFP4Tensor,MXTensor) all registeraten.view, which suggests this is an omission during the 0.18 rewrite rather thanan intentional design.
Because the dispatch tables are pure-Python per-class tables, this is not XPU-specific:
the same missing registrations should fail on CUDA as well (we only verified on XPU here).
Expected behavior
aten.viewonInt8Tensorandaten.abs/basic arithmetic onFloat8Tensorshouldeither be implemented (shape/view on the quantized data; abs/arith via dequantize
then re-quantize), or
quantize_()should explicitly no-op on already-quantizedmodules instead of failing per-module.
Workaround (verified working)
Runtime registration into the class dispatch tables works:
With these registrations, both
int8andfloat8quantize_()+ forward pass on XPU.