Repository navigation
Fix Int4OpaqueTensor A16W4 crash on mismatched activation dtype - #4969
Open
tritsystem wants to merge 1 commit into
Open
tritsystem wants to merge 1 commit into
tritsystem wants to merge 1 commit into
Conversation
Int4OpaqueTensor's docstring documents A16W4 (weight-only) support for
"float16/bfloat16/float32 activation + int4 weight", implying any of
these activation dtypes works regardless of the dtype the weight
happened to be in at quantization time. The forward implementation
passed the activation straight into
torch.ops.aten._weight_int4pack_mm_for_cpu without casting it to match
scale_and_zero.dtype. That aten op requires an exact dtype match (no
implicit promotion at this boundary, unlike plain elementwise ops), so
any mismatched (weight_dtype, activation_dtype) pair raised:
RuntimeError: expected scalar type X but found Y
Verified directly against Int4OpaqueTensor.from_hp() + F.linear() for
the full 3x3 dtype grid: only the 3 matching pairs worked before the
fix, all 6 mismatched pairs crashed. The sibling CUDA tensor class,
Int4TilePackedTo4dTensor, already casts its activation to its single
supported dtype (bfloat16) before calling its tinygemm op; this applies
the same cast-before-op pattern here, to scale_and_zero.dtype instead
of a hardcoded dtype, since Int4OpaqueTensor supports three.
Added test_linear_mismatched_activation_dtype, parametrized over all 9
(weight_dtype, act_dtype) combinations x both choose-qparams algorithms
(18 cases). Confirmed red/green: reverting only the source fix fails
exactly the 12 mismatched-dtype cases (6 pairs x 2 algorithms); the 6
matching-dtype cases pass either way. All 18 pass with the fix. Ran the
full existing test_int4_opaque_tensor.py suite before and after the fix:
36 pre-existing failures are present identically on unmodified main
(appear to be torch.compile-related issues in this environment, not
caused by or related to this change) with no new failures introduced.
Could not test the from_hp_da8w4 (DA8W4) path or any CUDA/XPU/NPU int4
tensor class, since this change only touches the CPU-only A16W4 path.
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
tritsystem
requested review from
Xia-Weiwen,
andrewor14,
jerryzh168 and
vkuzo
as code owners
October 3, 2026 05:36
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/ao/4969
Note: Links to docs will display an error until the docs builds have been completed. This comment was automatically generated by Dr. CI and updates every 15 minutes. |
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Int4OpaqueTensor's A16W4 (weight-only) docstring documents support for"float16/bfloat16/float32 activation + int4 weight", implying any of these
activation dtypes works regardless of the dtype the weight happened to be
in at quantization time. In practice, the
aten.linear.defaultforwardimplementation (
torchao/prototype/quantization/int4/int4_opaque_tensor.py)passed the activation straight into
torch.ops.aten._weight_int4pack_mm_for_cpuwithout casting it to matchscale_and_zero.dtype. That raw aten op requires an exact dtype match withno implicit promotion at this boundary (unlike plain elementwise ops on
current PyTorch, which silently type-promote), so any mismatched
(weight_dtype, activation_dtype)pair raises:Mechanism
scale_and_zerois fixed at whatever dtypeweightwas quantized from(
Int4OpaqueTensor.from_hp), e.g. float32 if thenn.Linearwasn't castbefore
quantize_()was called. If the caller later runs inference with adifferent float dtype (a very normal thing to do — e.g. quantize a model
as-loaded, then call it with a bfloat16 activation), the op call above
crashes with a cryptic dtype error that gives no hint the real cause is a
weight/activation dtype mismatch, contradicting the class's own documented
contract.
The sibling CUDA tensor class,
Int4TilePackedTo4dTensor, already handlesthis correctly by casting its activation to its one supported dtype
(bfloat16) before calling its own tinygemm op. This PR applies the same
cast-before-op pattern to
Int4OpaqueTensor, casting toscale_and_zero.dtype(since this class supports three activation dtypes, not just one).
Verification
Directly exercised
Int4OpaqueTensor.from_hp()+F.linear()for the full3x3
(weight_dtype, act_dtype)grid over{float32, bfloat16, float16}:mismatched pairs raise
RuntimeError: expected scalar type X but found Y.activation dtype.
Added
test_linear_mismatched_activation_dtypetotest/prototype/test_int4_opaque_tensor.py, parametrized over all 9(weight_dtype, act_dtype)pairs x bothuse_hqqsettings (18 cases).Red/green confirmed by reverting only the source fix: exactly the 12
mismatched-dtype cases fail (6 pairs x 2 algorithms); the 6 matching-dtype
cases pass regardless. With the fix, all 18 pass.
Ran the full existing
test_int4_opaque_tensor.pysuite before and afterthis change: 36 pre-existing failures are present identically on unmodified
main(appear to betorch.compile-related in this local environment —Python 3.14 / torch 2.13 CPU build — unrelated to this change) with no new
failures introduced by this fix.
Ran
ruff checkandruff formaton both changed files — clean.Verification limitation
Could not test the
from_hp_da8w4(DA8W4, int8-activation) path, or anyCUDA/XPU/NPU int4 tensor class, since this change only touches the
CPU-only A16W4 path (
_weight_int4pack_mm_for_cpu). This machine has noCUDA-enabled PyTorch build installed, so GPU-only code paths in this repo
could not be exercised at all.