[Common] Preserve NaN through half-precision MXFP8 amax reductions - #3574
CodeAlex52 wants to merge 3 commits into
Conversation
8821350 to
82b8d56
Compare
|
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>
31365cb to
46a5279
Compare
for more information, see https://pre-commit.ci
denera
left a comment
There was a problem hiding this comment.
Thanks for the PR!
Overall looks good, but there are two things that need to be addressed before we can test and merge:
test_cast_mxfp8.cucurrently 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.- 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
| 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); |
There was a problem hiding this comment.
| 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.
| }; | ||
|
|
||
| // ---- rowwise kernel (MXFP8 1D scaling, rowwise-only layout) ------------ | ||
| Tensor output_rowwise("output_rowwise", {rows, cols}, DType::kFloat8E4M3, |
There was a problem hiding this comment.
| 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.
| } | ||
|
|
||
| // ---- bidimensional kernel (rowwise + colwise layouts) ------------------ | ||
| Tensor output_bidim("output_bidim", {rows, cols}, DType::kFloat8E4M3, |
There was a problem hiding this comment.
| 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>(); |
There was a problem hiding this comment.
| 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.
| // 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). |
There was a problem hiding this comment.
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>
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 throughptx::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
maxwithout the.NaNqualifier 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
ptx::abs_max_nan_2xvariants (max.NaN.xorsign.abs.bf16x2/max.NaN.xorsign.abs.f16x2) — same element selection and sign semantics asabs_max_2x, except NaN propagates.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 helpermx_scale_reciprocal_x2keeps its non-NaN contract.abs_max_2xremains for its other callers).Testing
tests/cpp/operator/test_cast_mxfp8.cu: newCastMXFP8NaNScalingregression — 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.max.NaN.xorsign.abs.bf16x2in flashinfer's MXFP8 kernels,max.NaN.f32in cutlass) and the downstream NaN handling is verified at source level.Closes #3550