Skip to content

[Common, PyTorch] Add grouped NVFP4 dequantization - #3591

Open
harshithkantamneni wants to merge 5 commits into
NVIDIA:mainfrom
harshithkantamneni:harshith-group-dequantize-nvfp4
Open

harshithkantamneni wants to merge 5 commits into
NVIDIA:mainfrom
harshithkantamneni:harshith-group-dequantize-nvfp4

Conversation

@harshithkantamneni

@harshithkantamneni harshithkantamneni commented Sep 30, 2026 •

Copy link
Copy Markdown

Description

nvte_group_dequantize / tex.group_dequantize supported MXFP8 and FP8 block scaling but not NVFP4. This adds NVFP4 for rowwise data with compact (non-swizzled) E4M3 scales, as #2726 asks, and fixes the PyTorch binding so it passes an NVFP4 grouped tensor's element count, amax and swizzle flag to the kernel.

Fixes #2726

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

  • common/cast/nvfp4/group_dequantize_nvfp4.cuh (new): one thread per 16-element block, with the same arithmetic as dequantize_nvfp4.cuh, so output matches per-tensor dequantization bitwise. The owning tensor comes from row / rows_per_tensor (equal shapes) or common::find_tensor_from_offsets over tensor_offsets (varying first dimension, empty tensors included). tensor_offsets is read on the device, so there is no device-to-host copy and the call can be captured in a CUDA graph.
  • common/cast/dispatch/dequantize.cuh: route NVTE_NVFP4_1D_SCALING to the new kernel.
  • common/include/transformer_engine/cast.h: document NVFP4 support and its requirements.
  • pytorch/csrc/extensions/cast.cpp (group_dequantize), for NVFP4 only:
    • size the FP4 data in elements, not bytes, so CheckInputGroupedTensor accepts it;
    • pass amax and _with_gemm_swizzled_scales, which were not set before;
    • reject row-scaled NVFP4 and an E4M3 max other than 448, which the grouped tensor cannot describe.
  • tests/cpp/operator/test_dequantize_nvfp4_grouped.cu (new, registered in CMake): grouped vs. per-tensor nvte_dequantize, bitwise, over 6 shapes x 3 output types x shared/per-tensor amax.
  • tests/pytorch/test_grouped_tensor.py: grouped vs. per-tensor dequantize from tex.group_quantize output (including an empty tensor), the swizzled-scale error, and CUDA graph capture.

Scope

  • Rowwise data only. Compact scales only; swizzled scales raise an error.
  • Tensors must share the last dimension, and the last dimension and each tensor's first dimension must be multiples of 128, as the grouped quantizer (graph_safe_group_row_cast_col_hadamard_transform_cast_fusion.cu) requires. This makes the padded per-tensor scale layout one dense [rows, cols / 16] array. For equal shapes the per-tensor row count is checked on the host. With a varying first dimension the kernel reports a misaligned tensor through common::get_tensor_rows_num, once per tensor in the first block as group_quantize_mxfp8.cuh does (NVTE_DEVICE_ERROR, so it stops only in debug builds).
  • This is narrower than my plan on Dequantization support for the grouped tensor - NVFP4 #2726, where the three questions have no answer yet:
    1. Columnwise data is left out: per-tensor NVFP4 dequantize reads rowwise data only, and the grouped quantizer's columnwise data is Hadamard-transformed.
    2. VARYING_LAST_DIM and VARYING_BOTH_DIMS are rejected, because the grouped NVFP4 quantizer on main requires a constant last dimension. The kernel is a plain one and does not depend on [Common] Group NVFP4 Quantize Kernels  #3458.
    3. E4M3 scales only (UE5M3 is in Prototype NVFP4 with FP8 UE5M3 block scales #3325, see note 1).

Testing

Both commits were tested on one B300 (SM103, NVTE_CUDA_ARCHS=103a) at fb0fed6. The first commit was also tested on one B200 (SM100, 100a) at 1770259, with the same results. torch 2.8.0+cu129, cuDNN 9.26, CUDA 12.9. Current main does not build with CUDA 12.8: ptxas rejects the ld.global.nc.L2::evict_first fallback in util/ptx.cuh (from #3459) with "'.L1::eviction_priority' syntax expected", while docs/installation.rst lists 12.8+ for Blackwell.

  • test_operator --gtest_filter='*GroupedDequantizeNVFP4*': 33 passed, 3 skipped (the one-tensor per-tensor-amax cases, which duplicate the shared-amax ones), 0 failed. The reference, per-tensor nvte_dequantize, is itself checked against a CPU reference in test_dequantize_nvfp4.cu.

  • --gtest_filter='*Dequantize*': 777 passed, 219 skipped (including the 144 grouped FP8-blockwise cases, which run on SM90 only), 0 failed.

  • Row-count report: a grouped tensor with first dimensions 128 and 384 prints no message; one with 192 and 320 prints the get_tensor_rows_num message exactly twice, once per tensor, and in this release build the call continues, as for MXFP8.

  • New PyTorch tests: 8 passed. NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 pytest tests/pytorch/test_grouped_tensor.py: 167 passed, 71 skipped, 0 failed.

  • Each planted bug below was a temporary edit, rebuilt and run, then restored; all new tests pass again afterwards.

    Planted bug New tests failing (C++ of 33 / PyTorch of 8)
    amax always read from tensor 0 15 / 6
    row-to-tensor lookup off by one (row * cols - 1 passed to find_tensor_from_offsets) 9 / 6
    binding sizes the FP4 data in bytes 0 / 8
  • The pre-commit hooks (black 24.4.2, clang-format 18.1.6), cpplint 1.6.0 on the changed transformer_engine/ files, and qa/L0_license pass. No new build warnings in the changed files.

  • Not run: the full qa/L0_cppunittest and qa/L0_pytorch_unittest launchers, only the subsets above.

Notes for reviewers

  1. Prototype NVFP4 with FP8 UE5M3 block scales #3325 edits the same group_dequantize lines in cast.cpp and changes two things this PR relies on: the default nvfp4_e4m3_max becomes 0 (the scale dtype's max), which the nvfp4_e4m3_max == 448 check here would reject, and disable_second_level_scale allows NVFP4 without an amax, which this kernel rejects. Whichever lands second needs to handle both, and UE5M3 needs a scale-type template parameter in this kernel.
  2. If [Common] Group NVFP4 Quantize Kernels  #3458 changes the grouped NVFP4 scale layout, this kernel's indexing (one row of cols / 16 scales per data row) needs to follow it.
  3. tex.group_quantize for NVFP4 needs first_dims. Without it the grouped tensor has no tensor_offsets (build_grouped_tensor_offsets returns nullopt), and both graph-safe kernels read offsets[num_tensors] without a check (graph_safe_group_hadamard_transform.cu:241, graph_safe_group_row_cast_col_hadamard_transform_cast_fusion.cu:224), which fails with an illegal memory access. So the PyTorch tests pass first_dims, and the equal-shape path is covered by the C++ test.

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

AI usage: I used an AI assistant (Claude Code) to write code, tests, the validation scripts and this description. I reviewed every change, directed the tests, and ran the B200 validation.

nvte_group_dequantize had no NVFP4 path (NVIDIA#2726). Add one for rowwise
data with compact (non-swizzled) E4M3 scales, as written by the grouped
NVFP4 quantizer: one FP32 amax for the group or one per tensor, and
tensors that share the last dimension (equal shapes or a varying first
dimension). Each thread dequantizes one 16-element block with the same
arithmetic as the single-tensor NVFP4 kernel, so the results match it
bitwise.

Fix the PyTorch group_dequantize binding for NVFP4: size the FP4 data in
elements rather than bytes, pass the amax and the swizzled-scale flag,
and reject row-scaled NVFP4 and an E4M3 max other than 448, which the
grouped tensor cannot describe.

Add a C++ test that compares grouped against per-tensor dequantization
bitwise, and PyTorch tests against the grouped NVFP4 quantizer,
including an empty tensor, the swizzled-scale error and CUDA graph
capture.

Signed-off-by: Harshith Kantamneni <hkantamneni2@wisc.edu>
@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Sep 30, 2026
@harshithkantamneni
harshithkantamneni marked this pull request as ready for review September 30, 2026 00:22
@greptile-apps

greptile-apps Bot commented Sep 30, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 5/5

[Medium risk] Adds grouped NVFP4 dequantization support to the quantization library.

The PR appears safe to merge; the documentation change introduces no new issue, and no previous finding remains outstanding.

Summary

The PR adds grouped NVFP4 dequantization, updates the PyTorch binding and adds C++ and PyTorch coverage. Since the previous review, it clarifies the public API documentation to state that amax is optional when second-level scaling is disabled.

Reviews (5) · Last reviewed commit: "Document that grouped NVFP4 dequantizati..." · Reviewed by Greptile

Comment thread transformer_engine/common/cast/nvfp4/group_dequantize_nvfp4.cuh Outdated
@harshithkantamneni
harshithkantamneni marked this pull request as draft September 30, 2026 00:33
Look up the owning tensor with common::find_tensor_from_offsets instead
of a local binary search, and report per-tensor first dimensions that
are not multiples of 128 through common::get_tensor_rows_num, once per
tensor in the first block, as group_quantize_mxfp8.cuh does. The dense
scale indexing depends on that alignment, which was only checked on the
host for equal shapes.

Signed-off-by: Harshith Kantamneni <hkantamneni2@wisc.edu>
@harshithkantamneni
harshithkantamneni marked this pull request as ready for review September 30, 2026 01:50
Merge NVIDIA/TransformerEngine main at bba2d2b without rewriting the
existing PR commits. Restore the is_nvfp4 declaration removed by the
automatic merge with the scale_inv_dtype changes from NVIDIA#3325.

Accept nvfp4_e4m3_max=0 as the E4M3 default as well as explicit 448,
and cover both values in the grouped dequantization numerical test.
Keep the existing compact E4M3 scale and FP32 amax requirements.

Validation: compile the complete cast.cpp translation unit with CUDA
13.4 and PyTorch 2.13; run pre-commit on both changed paths, the C++
portion of L0_pytorch_lint, and L0_license. GPU numerical tests were
not run: the local SM120 GPU cannot execute grouped NVFP4 quantization,
which supports SM100-SM110.

Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Comment thread transformer_engine/common/cast/nvfp4/group_dequantize_nvfp4.cuh Outdated
group_quantize with disable_second_level_scale=True produces a grouped
NVFP4 tensor without an amax. Treat a missing amax as a unit global
scale, as dequantize_nvfp4.cuh does, instead of rejecting the tensor.
Cover it with a NO_AMAX mode in the C++ grouped test and with
disable_second_level_scale in the PyTorch grouped round trip.

Signed-off-by: Harshith Kantamneni <hkantamneni2@wisc.edu>
Comment thread transformer_engine/common/cast/nvfp4/group_dequantize_nvfp4.cuh
@harshithkantamneni

Copy link
Copy Markdown
Author

Thanks for merging main and restoring the is_nvfp4 declaration, @ptrendx.

The Greptile P1 on 8375386 was right: group_quantize with disable_second_level_scale=True returns a grouped NVFP4 tensor without an amax, and grouped dequantization rejected it. 227c3bd treats a missing amax as a global scale of 1, the same as dequantize_nvfp4.cuh, and covers it with a NO_AMAX mode in the C++ grouped test and with disable_second_level_scale in the PyTorch grouped round trip.

Your merge message mentions the GPU numerical tests could not run on SM120, so I ran them on a B200 (sm_100a) with CUDA 13.4, PyTorch 2.13.0+cu130 and cuDNN 9.27, on 227c3bd (your merge plus the fix):

Tests Passed Failed Skipped
C++ GroupedDequantizeNVFP4 51 0 3
C++ *Dequantize* 831 0 219
tests/pytorch/test_grouped_tensor.py with NVTE_GROUPED_LINEAR_SINGLE_PARAM=1, as in qa/L0_pytorch_unittest 185 0 71

To check that the new tests exercise the new path, I changed the default amax from 6 * 448 (global scale 1) to 1.0: exactly the 18 NO_AMAX C++ cases and the 12 disable_second_level_scale=True PyTorch cases failed, and everything passed again after restoring it.

Signed-off-by: Harshith Kantamneni <hkantamneni2@wisc.edu>

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

community-contribution PRs from external contributor outside the core maintainers, representing community-driven work.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Dequantization support for the grouped tensor - NVFP4

2 participants