Skip to content

Skip register_constant for Enum subclasses in register_as_pytree_constant - #4971

Open
vijay-kapse wants to merge 1 commit into
pytorch:mainfrom
vijay-kapse:fix/skip-pytree-constant-for-enums
Open

vijay-kapse wants to merge 1 commit into
pytorch:mainfrom
vijay-kapse:fix/skip-pytree-constant-for-enums

Conversation

@vijay-kapse

Copy link
Copy Markdown

Fixes #4848.

Problem

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 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:

class file
Enum ScaleCalculationMode prototype/mx_formats/config.py
Enum KernelPreference, a (str, Enum) quantization/quantize_/common/kernel_preference.py
non-Enum Float8TrainingOpConfig, MXFP8TrainingOpConfig prototype/moe_training/config.py

Importing kernel_preference on this branch versus main:

main:        W1003 ... <enum 'KernelPreference'> is an Enum subclass ...
this branch: (no pytree warning)

Tests

TestRegisterAsPytreeConstant patches torch.utils._pytree.register_constant and asserts it is
not called for a plain Enum or 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.

4 passed

Reverting only torchao/utils.py fails exactly the two Enum tests, while the non-Enum and
import cases keep passing, so the coverage discriminates:

FAILED test/test_utils.py::TestRegisterAsPytreeConstant::test_enum_is_not_registered
FAILED test/test_utils.py::TestRegisterAsPytreeConstant::test_str_enum_is_not_registered
2 failed, 2 passed

Full test/test_utils.py is green (16 passed). ruff format --check and ruff check are clean
on 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.

…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.
@pytorch-bot

pytorch-bot Bot commented Oct 3, 2026

Copy link
Copy Markdown

🔗 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.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Oct 3, 2026

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Avoid re-registering native Enum pytree constants

1 participant