Skip to content

[PyTorch] Introduce QuantizerRole - #2620

Merged
negvet merged 72 commits into
NVIDIA:mainfrom
negvet:semantic_quantizer_roles
May 11, 2026
Merged

negvet merged 72 commits into
NVIDIA:mainfrom
negvet:semantic_quantizer_roles

Conversation

@negvet

@negvet negvet commented Jan 23, 2026 •

Copy link
Copy Markdown
Collaborator

Description

Introducing QuantizerRole

@dataclasses.dataclass(frozen=True)
class QuantizerRole:
    module_type: str = ""   # e.g. "linear", "grouped_linear", "dpa"
    tensor_type: str = ""   # e.g. "input", "weight", "grad_output", "qkv", "s"
    name: str = ""          # instance name, e.g. "qkv", "proj", "fc1", "fc2"

This is an API that allows to go down to "set this LayerNormLinear in this transformer layer to be less aggressively quantized." (fine-grained, per-module/per-tensor quantization control mechanism)
See test_custom_recipe.py::test_custom_recipe_quantization_targets().

Quantizer factory uses roles to dispatch according to its needs.

TE module/op emits a list of QuantizerRole:

  • Linear, LayerNormLinear, LayerNormMLP emit module_type="linear" with tensor_type in {"input", "weight", "grad_output"}.
  • GroupedLinear emits module_type="grouped_linear".

CustomRecipe accepts a qfactory callable that receives QuantizerRole and returns a quantizer.

Factories can be composed - e.g., dispatch (to different sub-factories as an option) based on module_type (dpa vs linear) and then refine based on tensor_type.

Summary:

  • Modules implement get_quantizer_roles() that returns a list of QuantizerRole objects.
  • During set_meta_tensor(), modules call get_quantizer_roles() and pass roles to RecipeState.create().
  • RecipeState.create() assigns roles to the state (e.g., CustomRecipeState.roles).
  • CustomRecipeState.make_quantizers() calls qfactory(role) for each role to create quantizers.
  • The factory can inspect role.module_type, role.tensor_type, and role.name to dispatch to different quantizers.

This PR enables granular control over the recipe. Which might have some limitations and edge cases.

Type of change

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

Please list the changes introduced in this PR:

  • Change A
  • Change B

Checklist:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes

negvet and others added 4 commits January 23, 2026 15:14
…ipe state

Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
@negvet
negvet requested review from cyanguwa and timmoon10 January 23, 2026 15:32
@greptile-apps

greptile-apps Bot commented Jan 23, 2026 •

Copy link
Copy Markdown
Contributor

Greptile Summary

This PR introduces QuantizerRole — a frozen dataclass giving callers fine-grained, per-module/per-tensor control over the CustomRecipe quantizer factory. Modules now implement get_quantizer_roles(), and CustomRecipeState.make_quantizers() passes each role to the user factory. A new DelayedScalingRequest dataclass allows factories to request TE-managed stateful DS quantizers.

  • New API surface: QuantizerRole, QuantizerRequest, DelayedScalingRequest are exported; get_quantizer_roles() is added to TransformerEngineBaseModule, BasicOperation, and all concrete modules.
  • CustomRecipeState overhaul: The old hardcoded string-role dispatch is replaced with structured role objects; DelayedScalingRequest handling composes an inner DelayedScalingRecipeState for DS slots, including state preservation across role-driven rebuilds.
  • Boundary-role plumbing: MultiheadAttention._update_output_quantizer_roles wires consumer identity at the QKV→DPA and DPA→Proj boundaries on every forward pass.

Confidence Score: 3/5

The PR introduces significant new machinery; several defects in the changed code mean certain configurations will crash or produce silently wrong quantizers.

The asymmetric delayed-scaling crash in add_fp8_tensors_to_global_buffer is a latent runtime error that fires whenever a factory emits DelayedScalingRequest for one direction but not the other. The same has_delayed_scaling_state guard change also exposes restore_fp8_meta_tensors to a None.copy() crash. These are on top of assert-based length validations stripped by -O, the CustomRecipeState early-return that silently keeps stale quantizers when the recipe object changes, and the sequence-parallel amax check now being an assert instead of a hard error.

transformer_engine/pytorch/quantization.py (add_fp8_tensors_to_global_buffer and restore_fp8_meta_tensors), transformer_engine/pytorch/module/base.py (assert-based validation and CustomRecipeState early-return guard), transformer_engine/pytorch/ops/op.py (assert-based validation and recipe-identity check)

Important Files Changed

Filename Overview
transformer_engine/pytorch/quantization.py Core quantization state; adds QuantizerRole, DelayedScalingRequest, and CustomRecipeState overhaul. add_fp8_tensors_to_global_buffer can crash with TypeError when CustomRecipeState has DS active for only one direction.
transformer_engine/pytorch/module/base.py Adds get_quantizer_roles(), output/grad_input_quantizer_role properties, and inherit_state_from rebuild path. Uses assert for roles-length validation (skipped with -O) and CustomRecipeState early-return ignores recipe identity.
transformer_engine/pytorch/ops/op.py Adds get_quantizer_roles() and roles= parameter to RecipeState.create(). Roles-length validation uses assert (stripped with -O); type-identity check misses CustomRecipe instance swap.
transformer_engine/pytorch/attention/multi_head_attention.py Adds _update_output_quantizer_roles() to wire boundary roles between QKV, DPA and proj on every forward pass. Logic is correct; setter equality check prevents spurious re-initialization.
transformer_engine/pytorch/custom_recipes/quantization_recipes_base.py New reference factories for built-in recipes via CustomRecipe. current_scaling_quantizer_factory silently returns E4M3 for grad_input slots (should be E5M2 to match built-in Float8CurrentScaling).
transformer_engine/pytorch/custom_recipes/quantization_factory_examples.py New example factories for mixed-format recipes. Well-structured with clear dispatch logic; role=None handled defensively in all paths.
tests/pytorch/test_custom_recipe.py Comprehensive new tests for the QuantizerRole API including custom recipe quantization targets; factories are correctly updated to handle QuantizerRole objects instead of legacy role strings.
tests/pytorch/distributed/run_numerics_exact.py Updated factory handles role=None boundary slots and uses QuantizerRole field access. The legacy return-None branches are correctly replaced with real quantizer fallbacks.
transformer_engine/pytorch/module/layernorm_mlp.py Adds get_quantizer_roles() with fc1/fc2 named roles and internal boundary labeling. _warn_missing_output_quantizer_role is called with hardcoded fp8_grad=False, suppressing backward role warnings.
transformer_engine/pytorch/attention/dot_product_attention/dot_product_attention.py Adds get_quantizer_roles() for DPA's cuDNN slot layout, CustomRecipe early-return path in init_fp8_metadata, and state inheritance on rebuild. Role layout documentation is thorough.

Flowchart

%%{init: {'theme': 'neutral'}}%%
flowchart TD
    subgraph ModuleForward["Module forward (set_meta_tensor)"]
        A[fp8_meta_tensors_initialized?] -- No --> B[get_quantizer_roles]
        A -- Yes, CustomRecipeState exists --> EarlyReturn[Early return - stale if recipe changed]
        B --> C[RecipeState.create with roles]
        C --> D[inherit_state_from old state - preserves scale/amax_history]
        D --> E[make_quantizers]
    end

    subgraph CustomRecipeStateMQ["CustomRecipeState.make_quantizers"]
        E --> F[qfactory role for each slot]
        F --> G{Returns None?}
        G -- Yes --> Err1[ValueError: None rejected]
        G -- No --> H{DelayedScalingRequest?}
        H -- Yes --> I[_handle_delayed_scaling_requests - allocates inner DSRS]
        H -- No --> J[Use quantizer directly]
        I --> K[Splice Float8Quantizer into raw list]
    end

    subgraph GlobalBuffer["FP8GlobalStateManager.add_fp8_tensors_to_global_buffer"]
        L{_has_delayed_scaling_state?} -- True --> M[Loop fwd+bwd]
        M --> N{CustomRecipeState with DS?}
        N -- Yes --> O[Use inner_recipe key]
        N -- No --> P[Use outer recipe key]
        P --> Q[state.amax_history 0 - None if no DS - TypeError crash]
    end

    E --> L
Loading

Comments Outside Diff (1)

  1. transformer_engine/pytorch/quantization.py, line 519-540 (link)

    TypeError crash when DS is active for only one direction

    _has_delayed_scaling_state returns True as soon as either direction's CustomRecipeState has _has_delayed_scaling=True. When the other direction's CustomRecipeState has _has_delayed_scaling=False, the else branch at line 525 is taken and execution falls through to line 529 where fp8_meta[fp8_meta_tensor_key].amax_history[0] is evaluated. CustomRecipeState.amax_history returns None when _ds_state is absent, so None[0] raises TypeError.

    A practical trigger: a factory that returns DelayedScalingRequest only for forward DPA slots (e.g. "s") but no DS for any backward slot would leave scaling_fwd._has_delayed_scaling=True and scaling_bwd._has_delayed_scaling=False for that module. _has_delayed_scaling_state passes, and the loop crashes when it processes the backward direction. The same crash path exists in restore_fp8_meta_tensors (line 815) for the symmetric scenario where only bwd has DS.

    Add a continue guard for non-DS CustomRecipeState directions inside the loop to skip registration for directions without DS state.

Reviews (33): Last reviewed commit: "Merge branch 'main' into semantic_quanti..." | Re-trigger Greptile

@greptile-apps

This comment was marked as off-topic.

Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
greptile-apps[bot]

This comment was marked as outdated.

Signed-off-by: Evgeny <etsykunov@nvidia.com>
greptile-apps[bot]

This comment was marked as resolved.

Signed-off-by: Evgeny <etsykunov@nvidia.com>
greptile-apps[bot]

This comment was marked as resolved.

@timmoon10 timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Overall this design is quite clean and generalizable.

Comment thread transformer_engine/pytorch/quantization.py Outdated
Comment thread transformer_engine/pytorch/quantization.py Outdated
Comment thread transformer_engine/pytorch/quantization.py Outdated
Comment thread transformer_engine/pytorch/custom_recipes/quantization_nvfp4.py Outdated
Comment thread transformer_engine/pytorch/module/linear.py
Comment thread tests/pytorch/test_custom_recipe.py Outdated
negvet and others added 2 commits February 20, 2026 14:31
Signed-off-by: Evgeny Tsykunov <etsykunov@nvidia.com>
greptile-apps[bot]

This comment was marked as resolved.

negvet and others added 5 commits February 20, 2026 15:05
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
greptile-apps[bot]

This comment was marked as outdated.

Comment thread transformer_engine/pytorch/quantization.py Outdated
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
@negvet

negvet commented Apr 29, 2026

Copy link
Copy Markdown
Collaborator Author

Would the new quantizer setup be compatible with the existing one? Would TE tests (test_attention.py) or attention users need to update their code immediately?

This new setup is fully compatible with the current one - no need to immediately update anything. Legacy users will not notice it.
This new approach can be fully tested (and updated as required) before deprecating the current one.

@negvet

negvet commented Apr 29, 2026

Copy link
Copy Markdown
Collaborator Author

/te-ci pytorch L1

timmoon10
timmoon10 previously approved these changes Apr 29, 2026

@timmoon10 timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM, pending CI

negvet added 2 commits April 30, 2026 13:39
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
@negvet

negvet commented Apr 30, 2026

Copy link
Copy Markdown
Collaborator Author

/te-ci pytorch L1

timmoon10
timmoon10 previously approved these changes Apr 30, 2026
@timmoon10 timmoon10 mentioned this pull request May 5, 2026
8 of 13 tasks
negvet added 2 commits May 6, 2026 15:08
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
@negvet

negvet commented May 6, 2026

Copy link
Copy Markdown
Collaborator Author

QuantizerRole change triggers a recipe state rebuild (e.g. via output_quantizer_role setter).
This leads to the silent loss of the state of the stateful recipes, e.g. scale / amax_history for delayed scaling.

83c405b introduces RecipeState.inherit_state_from(), making buffers persistent. DelayedScalingRecipeState declares its buffers, stateless recipes have empty buffers (no op).

@negvet

negvet commented May 6, 2026

Copy link
Copy Markdown
Collaborator Author

/te-ci pytorch L1

timmoon10
timmoon10 previously approved these changes May 7, 2026
@ptrendx

ptrendx commented May 8, 2026

Copy link
Copy Markdown
Member

@negvet Could you fix the issues hit in the CI runs?

negvet added 2 commits May 8, 2026 13:53
Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@nvidia.com>
@negvet

negvet commented May 8, 2026

Copy link
Copy Markdown
Collaborator Author

/te-ci pytorch L1

@timmoon10 timmoon10 left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Test failures are unrelated

@negvet

negvet commented May 11, 2026

Copy link
Copy Markdown
Collaborator Author

/te-ci pytorch L1

@negvet

negvet commented May 11, 2026

Copy link
Copy Markdown
Collaborator Author

The same failures on the main

@negvet
negvet merged commit d73bfa1 into NVIDIA:main May 11, 2026
25 of 31 checks passed
faradawn pushed a commit to faradawn/TransformerEngine that referenced this pull request May 14, 2026
* Enable semantic roles emitted by module/op and comsumed by custom recipe state

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Update quantization factories

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Fix tests

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Swap tensor:module

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Better naming

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Introduce QuantizerRole frozen data class instead of a string

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Shrink module_type vocabulary

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Fix numerics exact test

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Set defaults, make custom recipe forward compatible

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* remove position from QuantizerRole

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Set good defaults

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Resolve naming: make every module/op distinguishable via name

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Configure output/grad_input roles, defaults to None

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Remove is_gemm()

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Enable base recipes via CustomRecipe and quantization factories

Signed-off-by: Evgeny <etsykunov@gmail.com>

* Add factory example - NVFP4 for Linear, MXFP8 for GroupedLinear

Signed-off-by: Evgeny <etsykunov@gmail.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Fix custom recipe test

Signed-off-by: Evgeny <etsykunov@gmail.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Test fine-grained quantization targets

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Add quantizer roles for attention (attn is wip)

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Enable statful recipes in the Custom recipe - Delayed Scaling support

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Fix save_original_input for custom delayed scaling

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Enable custom recipe for attn

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Make boundary role setting more explicit in MHA

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Make dpa role setting more intuitive

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Docstring for get_quantizer_roles() in base module

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Fix lint

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Restrict None roles

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Linter

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Minor fixes

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Test debug tools compat

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Fix pylint

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* fix test

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Fix lint

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Constructor takes roles kwarg + test fix

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Constructor takes roles kwarg + test fix (quantization.py)

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Fix attention: MXFP8, w/o CP

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Add test custom recipe

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Make Float8BlockScalingRecipeState and NVFP4BlockScalingRecipeState aware about QuantizerRole, dispatch on that + positional fallback if get_quantizer_roles() is not defined by the module/op

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Fix linter

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Fix CI

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Preserve delayed scaling state (buffers) when rebuild is triggered

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* Fix test, minor

Signed-off-by: Evgeny <etsykunov@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Fix distributed tests

Signed-off-by: Evgeny <etsykunov@nvidia.com>

---------

Signed-off-by: Evgeny <etsykunov@nvidia.com>
Signed-off-by: Evgeny Tsykunov <etsykunov@nvidia.com>
Signed-off-by: Evgeny <etsykunov@gmail.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Evgeny <etsykunov@gmail.com>
@negvet negvet changed the title [PyTorch] Introduce quantizer roles [PyTorch] Introduce QuantizerRole Aug 17, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants