Repository navigation
[diffusion] fix: do not interleave comfy-kitchen MXFP8 scales twice - #38028
Conversation
comfy-kitchen serializes MXFP8 block scales already in the SWIZZLE_32_4_4 order FlashInfer's cutlass/cutedsl kernels expect, while safetensors keeps their logical [N, K // 32] shape. SRT cannot tell them apart and runs block_scale_interleave over already-interleaved bytes, so the checkpoint loads and generates at full speed but decodes to noise. Route comfy-marked MXFP8 checkpoints through a Fp8LinearMethod subclass that flattens the serialized bytes instead. config.json-driven MXFP8 keeps SRT's behaviour: layer_markers is only set by resolve_comfy_checkpoint_quantization. Measured on one real MiniMax-H3 attention matrix: cosine 0.739481 -> 0.999641, relative L2 0.77703 -> 0.02679; the same fixed-seed 768p request goes from colour noise to a valid video with no measurable speed change. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
|
Per MAINTAINER.md, @mickqian @BBuf are the Diffusion oncalls for The 30-second version: The change is confined to Happy to adjust anything — in particular there is an open question in the description about |
Resolve the conflict with the online MXFP8 branch from sgl-project#37903: unserialized checkpoints still use MXFP8OnlineLinearMethod, serialized ones now get ComfyMXFP8LinearMethod. Adapt the override to the scale_u8 argument that sgl-project#40039 added to Fp8LinearMethod._process_mxfp8_linear_weight_scale. Only scales read from a comfy checkpoint skip the interleave; an explicit scale_u8 is converted from block-FP8 in row-major order and still goes through SRT. Update the unit test expectations accordingly. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
|
/tag-and-rerun-ci |
|
Thanks @mickqian for resolving the conflict with #37903 and carrying the override over to the CI on 65f8ce1 has finished. The new tests and the existing MXFP8 tests pass in
|
Motivation
comfy-kitchenserializes MXFP8 block scales already in theSWIZZLE_32_4_4byteorder that FlashInfer's cutlass / cutedsl kernels consume, while safetensors keeps
their logical
[N, K // 32]shape. Nothing in the file distinguishes those bytes fromrow-major scales, so
Fp8LinearMethod._process_mxfp8_linear_weight_scale(
python/sglang/srt/layers/quantization/fp8.py, cutlass/cutedsl branch) runsblock_scale_interleaveover bytes that are already interleaved.The second permutation is silent. No exception, no warning, no shape mismatch: the
checkpoint loads, the kernels run at full speed, and the sample decodes to dark,
low-amplitude colour noise. A user has no signal to distinguish this from a bad prompt
or a bad checkpoint.
Modifications
ComfyMXFP8LinearMethod, a thinFp8LinearMethodsubclass that hands FlashInferthe serialized scale bytes flattened instead of re-interleaving them.
MXFP8Config.get_quant_methodthrough it.Blast radius is limited to comfy checkpoints by construction.
layer_markersis onlyever populated by
resolve_comfy_checkpoint_quantization(
runtime/utils/quantization_utils.py), which is also the sole construction site ofMXFP8Configwith markers;MXFP8Config.from_config(theconfig.jsonpath) leaves itNone. The same predicate already gatescheckpoint_uses_native_qkv_layoutinMXFP8Config.__init__. For every non-comfy MXFP8 checkpoint, for every backendother than cutlass / cutedsl, and for an explicit
scale_u8(row-major scalesconverted from block-FP8, added in #40039), the call falls through to
super()unchanged.
Two unit tests in
python/sglang/multimodal_gen/test/unit/test_transformer_quant.py,both CPU-only and free of checkpoint fixtures:
test_comfy_mxfp8_scales_are_not_interleaved_twice— a comfy-marked config yieldsComfyMXFP8LinearMethodand its scale bytes survive load in order. It also assertsthat
Fp8LinearMethod._process_mxfp8_linear_weight_scalestill exists, so renamingthe SRT hook fails loudly instead of silently restoring the double interleave, and
that an explicit
scale_u8is still handed to SRT's interleave.test_config_driven_mxfp8_keeps_srt_scale_processing— aconfig.json-driven configstill delegates to SRT.
The existing
test_minimax_h3_global_mxfp8_metadata_selects_srt_per_layercontinues topass:
ComfyMXFP8LinearMethodis aFp8LinearMethod.Accuracy Tests
Measured on an RTX 5090 (sm_120) with a MiniMax-H3 Ref2VA MXFP8 single file
(
rzgar/...minimax_h3_ref2va_mxfp8.safetensors, 47.6 GB), one real attention weightmatrix, compared against the dequantized reference:
block_scale_interleaveapplied again (today)End to end on a fixed-seed 768p Ref2VA request: before the fix the run completed in
746.16 s and decoded to colour noise; after the fix the same request produces a valid
video. A locally converted INT8 checkpoint through the same pipeline was clean
throughout, which isolates the fault to the MXFP8 loading path rather than the prompt,
reference, text encoder, VAE or model code.
Speed Tests and Profiling
No measurable change. The edit is on the load path only, and it replaces a
block_scale_interleavecall with aview(-1). The fixed-seed request above ran746.16 s before and after within run-to-run noise; nothing in the forward path is
touched.
Checklist
pre-commititself could not be installed on this machine (no PyPIegress), so the configured hooks were run individually and all pass on both files:
isort --check-only(7.x and 8.x),ruff check --select=F401,F821,UP037,ruff format --diff, andblack --checkas a cross-check.multimodal_genunit test file next to theother MXFP8 tests. That file uses plain
unittest.TestCase;CustomTestCaseandregister_*_ciapply to registered suites undertest/registered/, not here.changes; this restores correct output for checkpoints that already load today.
mutation beyond what the base method already does, config class kept at the top of
the file.
Notes for reviewers
block_scale_interleave"may pad and/or reshape scales". Every Linear in the checkpointtested here has
N % 128 == 0, so flatten-only matches. If comfy-kitchen everserializes an unpadded scale block for a shape FlashInfer would pad, this path needs the
padding applied without the permutation. Happy to add a shape guard if you would prefer
to fail loudly there rather than assume.
Found while evaluating self-describing single-file MiniMax-H3 checkpoints on a single
RTX 5090 32 GB (sm_120, driver 616.56, Windows 11 + WSL2 + Docker, TP=1) against
20621aa14bda7726a8a968f326198eac61717fef. Since thenmaingained theMXFP8OnlineLinearMethodbranch (#37903) and thescale_u8hook argument (#40039);both are merged into this branch, and the cutlass/cutedsl branch in
srt/.../fp8.pystill re-interleaves
layer.weight_scale_invonmain, so the bug still reproducesthere.
🤖 Generated with Claude Code
CI States
Latest PR Test (Base): ❌ Run #37577160346
Latest PR Test (Extra): ❌ Run #37577160190
Latest PR Test (AMD ROCm 10): ⏳ Run #37577160383