Skip to content

[Common] Preserve NaN through half-precision MXFP8 amax reductions - #3574

Open
CodeAlex52 wants to merge 3 commits into
NVIDIA:mainfrom
CodeAlex52:fix/mxfp8-nan-amax-reduction
Open

CodeAlex52 wants to merge 3 commits into
NVIDIA:mainfrom
CodeAlex52:fix/mxfp8-nan-amax-reduction

Conversation

@CodeAlex52

Copy link
Copy Markdown

Summary

The half-precision MXFP8 scale paths lose NaN during the block amax reduction, so a block containing NaN is quantized with a finite E8M0 scale instead of the exceptional code 255, and the NaN payload semantics break. This makes the amax reduction NaN-preserving and adds mixed finite+NaN regression coverage for both the rowwise and bidimensional kernels.

Problem

A block containing a NaN element must take the exceptional E8M0 scale (float_to_e8m0(NaN) == 255) so the payload cast runs with a NaN scale (exp2f_rcp_2x(255) == NaN). On the half-precision rowwise and bidimensional paths the block amax is reduced through ptx::abs_max_2x (max.xorsign.abs.bf16x2 / max.xorsign.abs.f16x2), which ignores NaN: a NaN element combined against a finite extremum is dropped, and the block is quantized with a finite scale derived from the remaining finite values.

Root cause

PTX max without the .NaN qualifier returns the non-NaN operand when exactly one input is NaN. The later exceptional-value check never sees the lost NaN. Every downstream component is already prepared for a NaN amax (float_to_e8m0(NaN) == 0xFF, exp2f_rcp(255) == NaN, mx_scale_reciprocal(NaN) == kNaNReciprocal, and the colwise exceptional probe routes NaN columns to the per-half path); the reduction is the only missing link. This was flagged as the remaining correctness blocker during the review of #3459.

Fix

  • Add NaN-preserving ptx::abs_max_nan_2x variants (max.NaN.xorsign.abs.bf16x2 / max.NaN.xorsign.abs.f16x2) — same element selection and sign semantics as abs_max_2x, except NaN propagates.
  • Switch the MXFP8 rowwise (cast_rowwise.cuh: block_half_amax, the partner combine, pair_amax_to_float) and bidimensional (cast_bidim.cuh: rowwise/colwise passes, fold_pair_magnitude) reductions to the new variants. The guarded packed helper mx_scale_reciprocal_x2 keeps its non-NaN contract.
  • Non-NaN, integer, and other-kernel behavior is unchanged (the existing abs_max_2x remains for its other callers).

Testing

  • tests/cpp/operator/test_cast_mxfp8.cu: new CastMXFP8NaNScaling regression — a 64x32 bf16 tensor mixing large finite magnitudes with NaN lanes so that rowwise blocks (one row each) and colwise blocks (one column each) split between exceptional and finite scales; asserts the exceptional scale byte 255 for NaN-containing blocks, the finite reference scale for NaN-free blocks, and NaN payloads under the exceptional scale, for both the rowwise and bidimensional kernels. Runs on SM100 in CI.
  • Validation note: SM100 hardware is not available in external CI, so the RED state is documented by the issue reporter's hardware observation and the four-round review confirmation on [Common] Add non-TMA MXFP8 cast-only kernels for specialized rowwise-only and row+colwise  #3459; the instruction syntax and semantics are production-verified (max.NaN.xorsign.abs.bf16x2 in flashinfer's MXFP8 kernels, max.NaN.f32 in cutlass) and the downstream NaN handling is verified at source level.

Closes #3550

@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Sep 28, 2026
@CodeAlex52
CodeAlex52 force-pushed the fix/mxfp8-nan-amax-reduction branch 3 times, most recently from 8821350 to 82b8d56 Compare September 28, 2026 13:51
@greptile-apps

greptile-apps Bot commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 5/5

[Medium risk] Fixes NaN handling in half-precision floating-point reductions.

The PR appears safe to merge; no outstanding findings remain.

Summary

The PR makes half-precision MXFP8 amax reductions preserve NaN and adds coverage for exceptional scales and payloads in rowwise and bidimensional casts. The changes since the previous review correct the test’s host-buffer synchronization and make its tensor shapes explicit. All three previous findings are fixed.

Diagram

%%{init: {'theme': 'neutral'}}%%
flowchart LR
  A[Half-precision input block] --> B[NaN-preserving amax reduction]
  B --> C[E8M0 scale selection]
  C --> D[Rowwise or bidimensional FP8 cast]
  D --> E[Host-synchronized scale and payload assertions]
Loading

Reviews (5) · Last reviewed commit: "[Common] Fix MXFP8 NaN regression test b..."

Comment thread tests/cpp/operator/test_cast_mxfp8.cu Outdated
Comment thread tests/cpp/operator/test_cast_mxfp8.cu
Comment thread transformer_engine/common/cast/mxfp8/specialized/cast_bidim.cuh
The half-precision MXFP8 rowwise and bidimensional scale paths reduce
the block amax through ptx::abs_max_2x, i.e. max.xorsign.abs.bf16x2 /
max.xorsign.abs.f16x2.  These instructions ignore NaN: when a NaN
element is combined against a finite extremum the finite operand wins,
so the NaN is lost before the exceptional-value handling runs and the
block is quantized with a finite E8M0 scale instead of the exceptional
code 255.

Add NaN-preserving abs_max_nan_2x variants (max.NaN.xorsign.abs) and
use them in the MXFP8 scale paths so a NaN block reaches
float_to_e8m0(NaN) == 255 and the cast payload stays NaN.  The guarded
packed helper mx_scale_reciprocal_x2 keeps its non-NaN contract.
Non-NaN and integer behavior is unchanged.

Regression: mixed finite+NaN rowwise and bidimensional cases (128x256,
reaching the specialized kernels) asserting the exceptional scale byte,
NaN payloads, and finite scales for NaN-free blocks.

Closes NVIDIA#3550

Signed-off-by: CodeAlex <59381946+CodeAlex52@users.noreply.github.com>
@CodeAlex52
CodeAlex52 force-pushed the fix/mxfp8-nan-amax-reduction branch 3 times, most recently from 31365cb to 46a5279 Compare September 29, 2026 05:18
@denera denera self-assigned this Oct 2, 2026
@denera
denera self-requested a review October 2, 2026 16:40

@denera denera left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Thanks for the PR!

Overall looks good, but there are two things that need to be addressed before we can test and merge:

  1. test_cast_mxfp8.cu currently does not compile. I left in-line comments below on where it fails. These need to be addressed. Please compile and test them locally on your end.
  2. This PR changes the behavior for only these two BF16 kernels, but a NaN input still ends up with finite scales when it goes through other routes. This PR should preserve NaN for the following MXFP8 amax reduction cases as well:
    • Swizzled row+col kernel
    • Colwise only kernel
    • When cols % 128 != 0
    • FP16/FP32 inputs
    • The generic quantize kernels and the CUTEDSL backend

Comment thread tests/cpp/operator/test_cast_mxfp8.cu Outdated
constexpr size_t kRowwiseBlockCols = 128; // specialized rowwise kernel eligibility
constexpr size_t kRowBlockCols = 32; // MXFP8 block width along a row

Tensor input("input", {rows, cols}, DType::kBFloat16);

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Suggested change
Tensor input("input", {rows, cols}, DType::kBFloat16);
Tensor input("input", std::vector<size_t>{rows, cols}, DType::kBFloat16);

Ambiguous constructor, does not compile with NVCC 13.4 unless you narrow the type.

Comment thread tests/cpp/operator/test_cast_mxfp8.cu Outdated
};

// ---- rowwise kernel (MXFP8 1D scaling, rowwise-only layout) ------------
Tensor output_rowwise("output_rowwise", {rows, cols}, DType::kFloat8E4M3,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Suggested change
Tensor output_rowwise("output_rowwise", {rows, cols}, DType::kFloat8E4M3,
Tensor output_rowwise("output_rowwise", std::vector<size_t>{rows, cols}, DType::kFloat8E4M3,

Same issue as above, ambiguous constructor.

Comment thread tests/cpp/operator/test_cast_mxfp8.cu Outdated
}

// ---- bidimensional kernel (rowwise + colwise layouts) ------------------
Tensor output_bidim("output_bidim", {rows, cols}, DType::kFloat8E4M3,

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Suggested change
Tensor output_bidim("output_bidim", {rows, cols}, DType::kFloat8E4M3,
Tensor output_bidim("output_bidim", std::vector<size_t>{rows, cols}, DType::kFloat8E4M3,

Same issue as above, ambiguous constructor.

ASSERT_EQ(cudaGetLastError(), cudaSuccess);

const fp8e8m0 *scales_rowwise = output_rowwise.rowwise_cpu_scale_inv_ptr<fp8e8m0>();
const OutputType *out_rowwise = output_rowwise.rowwise_cpu_dptr<OutputType>();

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Suggested change
const OutputType *out_rowwise = output_rowwise.rowwise_cpu_dptr<OutputType>();
const OutputType *out_rowwise = output_rowwise.to_cpu().rowwise_cpu_dptr<OutputType>();

.rowwise_cpu_dptr() returns the host mirror without copying from device, so you have to have a .to_cpu() first to make sure the D2H copy actually triggers.

Comment on lines +947 to +954
// Regression test for https://github.com/NVIDIA/TransformerEngine/issues/3550:
// the amax reduction of the half-precision MXFP8 rowwise/bidimensional kernels
// used max.xorsign.abs (NaN-ignoring), so a NaN input element was dropped
// before the exceptional-value handling and its block was quantized with a
// finite scale. Every 32-element block containing NaN must take the
// exceptional E8M0 scale (255) and the cast payload must preserve NaN.
// The shape reaches the specialized kernels: 256 % 128 == 0 (rowwise) and
// 256 % 256 == 0 (bidimensional).

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Duplicate comment block here, same as the block above TEST(...). Please consolidate.

- Disambiguate the shape argument of the three Tensor constructions in the
  NaN scaling regression test; the braced list matched both the NVTEShape and
  the std::vector<size_t> overload, which does not compile with NVCC 13.4.
- Call to_cpu() before reading the rowwise/columnwise data pointers of the
  quantized outputs.  Unlike rowwise_cpu_scale_inv_ptr(), those accessors
  return the host mirror without copying, so the D2H copy has to be requested
  explicitly.
- Drop the duplicated comment block above the TEST, keeping the fuller one
  inside the test body.

Signed-off-by: CodeAlex52 <59381946+CodeAlex52@users.noreply.github.com>

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.

Preserve NaN through half-precision MXFP8 amax reductions

2 participants