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 2 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 2 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 adds an SM100 TMA-tiled kernel that computes rowwise and columnwise NVFP4 amax metadata in one pass, with runtime fallback and an environment kill switch.

  • Routes eligible forward and backward row-scaled quantization through the fused implementation.
  • Retains the two-kernel path for unsupported shapes, data types, output layouts, or disabled fusion.
  • Introduces one state-preservation regression for no-op quantization calls.

Confidence Score: 4/5

The no-op state-preservation regression should be fixed before merging because skipped quantization calls can erase both amax buffers.

The fused wrapper unconditionally queues zeroing of the amax outputs, while its kernel subsequently returns on the noop flag, leaving previously preserved quantization metadata cleared.

Files Needing Attention: transformer_engine/common/cast/nvfp4/quantize_transpose_nvfp4.cuh

Important Files Changed

Filename Overview
transformer_engine/common/cast/dispatch/quantize.cuh Adds symmetric forward and backward selection of the fused amax path while preserving the original fallback.
transformer_engine/common/cast/nvfp4/quantize_transpose_nvfp4.cuh Implements and gates the fused TMA reduction, but clears amax metadata before honoring the existing no-op contract.

Flowchart

%%{init: {'theme': 'neutral'}}%%
flowchart TD
  Q[Row-scaled NVFP4 quantization] --> C{Columnwise output requested?}
  C -- No --> R[Standalone rowwise amax]
  C -- Yes --> S{BF16, 128-aligned, buffers present, fusion enabled?}
  S -- No --> F[Standalone rowwise and columnwise amax]
  S -- Yes --> Z[Clear both amax buffers]
  Z --> K[Fused TMA row and column reduction]
  K --> N{noop equals 1?}
  N -- Yes --> E[Kernel returns; cleared metadata remains]
  N -- No --> A[Atomically publish row and column maxima]
Loading

Reviews (1): Last reviewed commit: "[pre-commit.ci] auto fixes from pre-comm..." | Re-trigger Greptile

Comment on lines +434 to +442
NVTE_CHECK(output->amax.numel() == rows, "Fused rowwise amax must have ", rows,
" entries, got ", output->amax.shape, ".");
NVTE_CHECK_CUDA(cudaMemsetAsync(row_amax_ptr, 0, rows * sizeof(float), stream));
}
if (do_col) {
NVTE_CHECK(output->columnwise_amax.numel() == cols, "Fused columnwise amax must have ", cols,
" entries, got ", output->columnwise_amax.shape, ".");
NVTE_CHECK_CUDA(cudaMemsetAsync(col_amax_ptr, 0, cols * sizeof(float), stream));
}

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

P1 No-op clears amax state

When an eligible fused quantization call has noop[0] == 1, this wrapper clears both amax buffers before the kernel returns without writing, causing skipped graph-replay or update calls to destroy the previously preserved quantization metadata.

Knowledge Base Used: Native GEMM and quantization kernels

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