Repository navigation
[diffusion] Fix native FP8 format handling for FLUX 3 rowwise linears - #42255
Conversation
|
/tag-and-rerun-ci |
|
Native AMD validation is now available for this exact head (
Main CUDA CI and XPU diffusion CI pass. AMD 1-GPU partitions 2 and 3 also pass. This is still not an all-green AMD run: the separate FLUX Action full request continues to contain non-finite values, JoyEcho output fails, H3 encounters its existing HIP allowlist/media admission errors, and LTX-2.5 audio VAE access returns HTTP 403. NPU again fails installing the Triton Ascend wheel (403). The targeted FP8 fix is now supported by real AMD execution as well as CUDA; no new change or relaxed tolerance was needed. Could the independent policy/fallback PRs be authorized for CI next, and the remaining platform/artifact failures be assessed separately for merge eligibility? |
Motivation
The MI300 diffusion unit run on #42121 fails the FLUX 3 FP8r checkpoint test with
Expected b.dtype() == at::kFloat8_e4m3fnuz, got Float8_e4m3fn(job). The rowwise linear currently passes serialized E4M3FN weights and activations directly to_scaled_mm, although gfx94 requires FNUZ. ModelOpt already normalizes that format, but two of its regression tests still assume FN and unadjusted scales.Modifications
_scaled_mm, reusing the existing FN-to-FNUZ helper. Both scales are adjusted; negative zero is normalized to avoid NaN. Clone checkpoint weights before the helper's in-place byte normalization.Accuracy Tests
Tested on an isolated RTX 5090 D v2 host, Python 3.12 / PyTorch 2.13.0+cu130, based on upstream
f6fcda8.test_flux3_action.py,test_modelopt_fp8_layerwise_offload_load.py,test_transformer_quant.py: 107 passed, 24 subtests passed. Includes the real CUDA FP8_scaled_mmcheckpoint forward.Please authorize AMD CI to verify the actual MI300 GEMM/offload path. This fixes the concrete format mismatch; it does not yet establish that the separate full FLUX Action non-finite-output failure is resolved. No thresholds or tolerances were relaxed.
Speed Tests and Profiling
Compatibility fix; no throughput claim. CUDA uses its existing FP8 path. The FNUZ path adds format normalization to make the previously invalid GEMM operands usable.
Checklist
CI States
Latest PR Test (Base): ✅ Run #37052849608
Latest PR Test (Extra): ❌ Run #37052849224
Latest PR Test (AMD ROCm 10): ❌ Run #37052849660