Skip to content

[Commo] Optimize grouped MXFP8 requantization - #3560

Open
Oleg-Goncharov wants to merge 26 commits into
NVIDIA:mainfrom
Oleg-Goncharov:pr_requantize_mxfp8_perf
Open

Oleg-Goncharov wants to merge 26 commits into
NVIDIA:mainfrom
Oleg-Goncharov:pr_requantize_mxfp8_perf

Conversation

@Oleg-Goncharov

Copy link
Copy Markdown
Collaborator

Description

Add a fused grouped MXFP8 requantization path for Blackwell GPUs. The new operation converts compact rowwise MXFP8 input into columnwise E4M3 MXFP8 without materializing an intermediate high-precision tensor in global memory.

This reduces memory traffic compared with separate grouped dequantization and quantization calls. The kernel uses a tiled TMA pipeline and packed data movement/conversion operations, and supports both compact and GEMM-swizzled output scaling factors. The quantization configuration selects either an FP32 intermediate or a faster BF16 intermediate.

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:

  • Add the nvte_group_requantize_mxfp8 C API for fused grouped rowwise-to-columnwise MXFP8 requantization.
  • Add an SM100+ TMA-based kernel supporting grouped tensors, FP8 input types, E4M3 output, and compact or GEMM-swizzled output scales.
  • Add optimized BF16 fast-math and FP32 intermediate-precision paths.
  • Optimize packed data access and conversion in the requantization kernel.
  • Add parameterized C++ correctness tests comparing the fused operation against separate dequantization and quantization.

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

Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
@greptile-apps

greptile-apps Bot commented Sep 22, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 2/5

[High risk] Refactors grouped MXFP8 requantization kernel and test suite.

The PR does not appear safe to merge while mixed-format NVFP4 GEMM fails and the outstanding MXFP8 output-store issue remains.

Findings

  1. P1 Security BF16 output can overflow ▶
  2. P2 Kernel variants lack coverage ▶
  3. P2 Test input is overwritten ▶

Summary

The PR adds fused grouped MXFP8 requantization with CUDA and optional CuTeDSL implementations, updates its grouped-tensor API and tests, and includes subsequent NVFP4 scale-format and GEMM changes.

  • The new NVFP4 scale dispatch does not handle mixed UE5M3/E4M3 operands.
  • Earlier unresolved findings remain relevant to the MXFP8 kernel’s partial BF16 store and the correctness test’s aliased input.

Reviews (8) · Last reviewed commit: "Merge branch 'main' into pr_requantize_m..."

Comment thread transformer_engine/common/cast/dispatch/quantize.cuh Outdated
Comment thread transformer_engine/common/cast/dispatch/quantize.cuh Outdated
Comment thread transformer_engine/common/cast/dispatch/gated.cuh Outdated
Comment thread transformer_engine/common/cast/dispatch/dequantize.cuh Outdated
Comment thread tests/cpp/CMakeLists.txt Outdated
Comment thread tests/cpp/operator/CMakeLists.txt Outdated
Comment thread tests/cpp/operator/test_requantize_mxfp8_grouped.cu Outdated
Signed-off-by: Oleg Goncharov <ogoncharov@nvidia.com>
@ptrendx

ptrendx commented Sep 23, 2026

Copy link
Copy Markdown
Member

@Oleg-Goncharov This PR seems to have some unrelated changes (like the nccl ep submodule). Could you fix that?

Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
@ptrendx ptrendx self-assigned this Sep 30, 2026
Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
@ptrendx ptrendx added the 2.21 label Sep 30, 2026
Adapt the cuDNN warp-specialized BF16 kernel to compact input scales,
E4M3/E5M2 input, TE output layouts, optional rowwise scale swizzling,
and device-side grouped offsets. Cache shape specializations in the
common C++ TVM FFI dispatcher and preserve CUDA fallbacks.

Add independent numerical/capacity tests, reproducible benchmarks and
profiling commands. Validate 55 direct cases, 52 PyTorch integration
cases and 146 common C++ cases; boundary/extreme memcheck passes.

On one GB200, improve all nine measured shapes over the grouped CUDA
kernel. Large shapes reach 5.9-6.4 TB/s effective logical bandwidth;
Nsight Compute measures 6.15 TB/s physical DRAM traffic. Column-only
performance matches upstream cuDNN; the small extra cost with both TE
scale outputs is explained by the added rowwise scale output.

Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Measure the native CUDA and cached CuTeDSL paths through the same
public grouped C API, using a C++ timing loop to exclude Python and
ctypes overhead. Drain queues outside timing, report bounded async
batches and graph capture separately, and calibrate thread CPU clocks.

Also measure the PyTorch binding with prequantized grouped inputs,
cached device offsets and output allocation included. Alternate backend
order across nine batches of 1000 calls and verify identical outputs.

On a GB200/Grace host, CuTeDSL adds about 0.6-0.8 us per preallocated
batched C API dispatch and 3-4 us per full PyTorch binding call. Record
all batch statistics, methodology, environment and reproduction commands.

Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
@ptrendx
ptrendx requested a review from ksivaman as a code owner October 1, 2026 18:59
ptrendx and others added 8 commits October 1, 2026 12:00
Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Expose one grouped C API with an optional BF16 dequantized output. Select the faster columnwise requantization kernel when that output is absent, and retain the fused materialization kernel when it is requested. Consolidate the PyTorch binding and preserve grouped shape, scale-layout, capacity-tail and offsets-only coverage.

Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Emit optional BF16 values directly from decoded registers without extra
shared memory. Specialize cached TVM-FFI functions on the output option and
retain native CUDA fallback for unsupported configurations. Use narrow
tiles when the doubled tile count still fits in one SM wave.

Add a preallocated C API benchmark that checks all outputs, identifies
both kernels, and separates compilation from CUDA graph GPU timing.
On one GB200, median BF16-output times across nine 100-call batches improve
from 3.27 to 2.97 us at 1024x256, 8.38 to 6.07 us at 8192x1024, and
111.52 to 90.29 us at 16384x8192.

Validate 146 common requantization tests, 171 focused PyTorch tests,
and all 72 common materialization cases after tile tuning. Check both
input FP8 formats, scale layouts, optional BF16 output, empty groups,
capacity tails, and CUDA fallback.

Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
Comment thread transformer_engine/common/cast/mxfp8/requantize_mxfp8.cu
Comment thread transformer_engine/common/cast/mxfp8/requantize_mxfp8.cu
Select the scaled FP8-to-BF16 instruction only when the CuTeDSL compiler supports CUDA 13.2. Otherwise, use the native-equivalent FP16 conversion and BF16 scale multiplication sequence. Exercise automatic and portable decoding for both FP8 formats, including every FP8 encoding at extreme scales.

Validation: 235 direct requantization tests passed with each CUDA 12.9 and CUDA 13.4 compiler on GB200. Pre-commit and focused pylint passed.
Signed-off-by: Przemek Tredak <ptredak@nvidia.com>
@ptrendx
ptrendx force-pushed the pr_requantize_mxfp8_perf branch from de60011 to b18d27e Compare October 2, 2026 23:08
Comment thread tests/cpp/operator/test_requantize_mxfp8_grouped.cu

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants