Repository navigation
Skip register_constant for Enum subclasses in register_as_pytree_constant - #4971
Open
vijay-kapse wants to merge 1 commit into
Open
vijay-kapse wants to merge 1 commit into
vijay-kapse wants to merge 1 commit into
Conversation
…tant
PyTorch handles Enum values as opaque pytree constants natively now, and calling
register_constant() on an Enum subclass is deprecated there. Importing any
TorchAO module that decorates one emits:
torch/utils/_pytree.py:630] <enum 'KernelPreference'> is an Enum subclass and
is now natively supported by torch.compile as an opaque value type. Calling
register_constant() on Enum subclasses is deprecated and will be an error in a
future release.
The decorator registered unconditionally:
def register_as_pytree_constant(cls):
torch.utils._pytree.register_constant(cls)
return cls
so the warning fires at import time, once per decorated Enum, in every process
that imports TorchAO. pytorch#4848 reports it repeating for each vLLM worker at startup.
It is also a forward-compatibility problem rather than only noise, since PyTorch
documents it as becoming an error.
The decorator is applied to both shapes, so the Enum case is the only one skipped:
Enum ScaleCalculationMode (prototype/mx_formats/config.py)
KernelPreference, a (str, Enum)
(quantization/quantize_/common/kernel_preference.py)
non-Enum Float8TrainingOpConfig, MXFP8TrainingOpConfig
(prototype/moe_training/config.py)
Before and after, importing kernel_preference:
without this change: W1003 ... <enum 'KernelPreference'> is an Enum subclass ...
with this change: (no pytree warning)
Adds TestRegisterAsPytreeConstant, which patches
torch.utils._pytree.register_constant and asserts it is not called for a plain
Enum or for a (str, Enum), is called exactly once for an ordinary class, and that
the decorator still returns the class in each case. Reverting the change to
torchao/utils.py fails exactly the two Enum tests.
Fixes pytorch#4848. The fix matches the approach described there.
vijay-kapse
requested review from
andrewor14,
jerryzh168 and
vkuzo
as code owners
October 3, 2026 17:57
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/ao/4971
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.
Fixes #4848.
Problem
PyTorch handles
Enumvalues as opaque pytree constants natively now, and callingregister_constant()on anEnumsubclass is deprecated there. Importing any TorchAO modulethat decorates one emits:
The decorator registered unconditionally:
so it fires at import time, once per decorated Enum, in every process that imports TorchAO.
#4848 reports it repeating for each vLLM worker at startup. It is also a forward-compatibility
problem rather than only noise, since PyTorch documents it as becoming an error.
Change
Only the Enum case is skipped. The decorator is applied to both shapes in the tree, and the
non-Enum configs still need the explicit registration:
ScaleCalculationModeprototype/mx_formats/config.pyKernelPreference, a(str, Enum)quantization/quantize_/common/kernel_preference.pyFloat8TrainingOpConfig,MXFP8TrainingOpConfigprototype/moe_training/config.pyImporting
kernel_preferenceon this branch versusmain:Tests
TestRegisterAsPytreeConstantpatchestorch.utils._pytree.register_constantand asserts it isnot called for a plain
Enumor for a(str, Enum)— the shape actually seen in practice —is called exactly once for an ordinary class, and that the decorator still returns the class
it decorated in each case.
Reverting only
torchao/utils.pyfails exactly the two Enum tests, while the non-Enum andimport cases keep passing, so the coverage discriminates:
Full
test/test_utils.pyis green (16 passed).ruff format --checkandruff checkare cleanon both files.
Credit
@malaiwah described this fix and linked a commit on their fork in #4848 five weeks ago. No PR
materialised, so I have written it up here — happy to close this in favour of theirs if they
would rather carry it.