Skip to content

[Common] row-scaled nvfp4 path: fuse row/col amax into a single TMA-tiled kernel - #3454

Open
cael-ling wants to merge 3 commits into
NVIDIA:mainfrom
cael-ling:pr/nvfp4-row-scaled-fused-amax
Open

[Common] row-scaled nvfp4 path: fuse row/col amax into a single TMA-tiled kernel#3454
cael-ling wants to merge 3 commits into
NVIDIA:mainfrom
cael-ling:pr/nvfp4-row-scaled-fused-amax

Conversation

@cael-ling

@cael-ling cael-ling commented Sep 1, 2026

Copy link
Copy Markdown
Contributor

Description

The NVFP4 row-scaled path that was originally proposed in #2931 computes per-row and per-column amax with two separate kernels. This PR fuses both directions into a single kernel that streams 128x128 chunks through shared memory via TMA (coalesced loads) and does the column reduction from SMEM. Amax is an exact max reduction, so results are byte-identical to the two-kernel path. The fused path is used only when the quantize call needs both directions (rowwise + columnwise amax) on a BF16 input with 128-aligned dims; the kernel then produces both amaxes in one pass. Any other case keeps the original two kernels. NVTE_NVFP4_FUSED_AMAX=0 forces the fallback at runtime.

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

  • Add compute_fused_amax_kernel (rowwise + columnwise) and its host wrappers
    fused_amax_supported / compute_fused_amax in quantize_transpose_nvfp4.cuh.
  • Dispatch to the fused kernel in dispatch/quantize.cuh (fwd and bwd) when supported, else fall back to the standalone amax kernels.
  • Add NVTE_NVFP4_FUSED_AMAX kill switch (default on).

Performance

Full-quantize median latency, fused (fused amax + cast) vs (row-wise & columnwise amax + cast), the cast kernel is byte-identical so the delta is the amax step:

shape fused (ms) fallback (ms) speedup
4096x4096 0.098 0.138 1.41x
8192x8192 0.148 0.360 2.44x
8192x16384 0.210 0.648 3.09x
32768x8192 0.346 1.222 3.53x
16384x16384 0.348 1.263 3.63x

Reproduce

Single Blackwell (SM100) GPU.

  • pytest tests/pytorch/nvfp4/test_nvfp4_quantize_exact.py -k "row_scaled and both_directions"
    passes; rowwise/columnwise amax and qx/qx_t match the reference exactly.
  • Fixed-input repeat runs give byte-identical amax (deterministic) and match the two-kernel path.
  • Add NVTE_NVFP4_FUSED_AMAX env var (default enabled): set to 0 to disable the fused path at runtime and fall back to the two standalone amax kernels, no rebuild needed.

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

The row-scaled path ran two amax kernels; the columnwise one read global
memory column-major (uncoalesced) and dominated runtime. Compute both
directions in one kernel that streams 128x128 chunks through shared memory
via TMA and reduces columns from SMEM.

Gated by fused_amax_supported (BF16, 128-aligned dims); other cases keep the
two-kernel path. Set NVTE_NVFP4_FUSED_AMAX=0 to force fallback.

Signed-off-by: Cael Ling <caell@nvidia.com>
@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Sep 1, 2026
@greptile-apps

greptile-apps Bot commented Sep 1, 2026

Copy link
Copy Markdown
Contributor

Greptile Summary

The PR optimizes row-scaled NVFP4 quantization by replacing separate rowwise and columnwise amax passes with a fused TMA-tiled kernel when the input and output satisfy the supported BF16 and alignment constraints.

  • Adds a fused row/column amax kernel for Blackwell-class GPUs.
  • Dispatches the fused implementation from forward and backward quantization paths.
  • Preserves the existing kernels as a fallback and provides the NVTE_NVFP4_FUSED_AMAX runtime kill switch.
  • The previous no-op state-preservation finding is resolved in the current implementation.

Confidence Score: 5/5

The PR appears safe to merge, with no outstanding findings or newly introduced issues since the previous review.

The previous no-op metadata-preservation finding is resolved, and no code has changed since that review; the fused path now leaves both amax buffers untouched when the no-op flag is set.

Important Files Changed

Filename Overview
transformer_engine/common/cast/dispatch/quantize.cuh Selects the fused amax implementation in forward and backward row-scaled quantization while retaining the existing fallback.
transformer_engine/common/cast/nvfp4/quantize_transpose_nvfp4.cuh Implements the fused TMA-tiled amax kernel, output initialization, support checks, launch wrapper, and runtime kill switch.

Flowchart

%%{init: {'theme': 'neutral'}}%%
flowchart TD
  A[Row-scaled NVFP4 quantization] --> B{Columnwise output requested?}
  B -- No --> E[Standalone rowwise amax]
  B -- Yes --> C{Fused path supported and enabled?}
  C -- Yes --> D[Fused TMA row and column amax]
  C -- No --> F[Standalone rowwise and columnwise amax]
  D --> G[NVFP4 cast and transpose]
  E --> G
  F --> G
Loading

Reviews (3): Last reviewed commit: "[Common] Preserve NVFP4 fused amax buffe..." | Re-trigger Greptile

Comment thread transformer_engine/common/cast/nvfp4/quantize_transpose_nvfp4.cuh
The fused wrapper zeroed both amax buffers with an unconditional memset while
the kernel returns early on noop[0]==1, so a skipped (graph-replay) call cleared
the previously published amax. Replace the memset with a noop-aware zero kernel
that returns early on the same flag, matching the standalone kernels' contract.

Signed-off-by: Cael Ling <caell@nvidia.com>
@cael-ling
cael-ling marked this pull request as draft September 7, 2026 06:43
@cael-ling
cael-ling marked this pull request as ready for review September 7, 2026 07:03
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.

1 participant