Skip to content

[diffusion] fix: do not interleave comfy-kitchen MXFP8 scales twice - #38028

Merged
mickqian merged 6 commits into
sgl-project:mainfrom
TyGu888:diffusion-mxfp8-comfy-scale
Oct 7, 2026
Merged

mickqian merged 6 commits into
sgl-project:mainfrom
TyGu888:diffusion-mxfp8-comfy-scale

Conversation

@TyGu888

@TyGu888 TyGu888 commented Sep 4, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

comfy-kitchen serializes MXFP8 block scales already in the SWIZZLE_32_4_4 byte
order
that FlashInfer's cutlass / cutedsl kernels consume, while safetensors keeps
their logical [N, K // 32] shape. Nothing in the file distinguishes those bytes from
row-major scales, so Fp8LinearMethod._process_mxfp8_linear_weight_scale
(python/sglang/srt/layers/quantization/fp8.py, cutlass/cutedsl branch) runs
block_scale_interleave over 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

  • Add ComfyMXFP8LinearMethod, a thin Fp8LinearMethod subclass that hands FlashInfer
    the serialized scale bytes flattened instead of re-interleaving them.
  • Route MXFP8Config.get_quant_method through it.

Blast radius is limited to comfy checkpoints by construction. layer_markers is only
ever populated by resolve_comfy_checkpoint_quantization
(runtime/utils/quantization_utils.py), which is also the sole construction site of
MXFP8Config with markers; MXFP8Config.from_config (the config.json path) leaves it
None. The same predicate already gates checkpoint_uses_native_qkv_layout in
MXFP8Config.__init__. For every non-comfy MXFP8 checkpoint, for every backend
other than cutlass / cutedsl, and for an explicit scale_u8 (row-major scales
converted 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 yields
    ComfyMXFP8LinearMethod and its scale bytes survive load in order. It also asserts
    that Fp8LinearMethod._process_mxfp8_linear_weight_scale still exists, so renaming
    the SRT hook fails loudly instead of silently restoring the double interleave, and
    that an explicit scale_u8 is still handed to SRT's interleave.
  • test_config_driven_mxfp8_keeps_srt_scale_processing — a config.json-driven config
    still delegates to SRT.

The existing test_minimax_h3_global_mxfp8_metadata_selects_srt_per_layer continues to
pass: ComfyMXFP8LinearMethod is a Fp8LinearMethod.

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 weight
matrix, compared against the dequantized reference:

scale handling cosine relative L2
checkpoint bytes flattened (this PR) 0.999641 0.02679
block_scale_interleave applied again (today) 0.739481 0.77703

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_interleave call with a view(-1). The fixed-seed request above ran
746.16 s before and after within run-to-run noise; nothing in the forward path is
touched.

Checklist

  • Formatted. pre-commit itself could not be installed on this machine (no PyPI
    egress), 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, and black --check as a cross-check.
  • Unit tests added, in the existing multimodal_gen unit test file next to the
    other MXFP8 tests. That file uses plain unittest.TestCase; CustomTestCase and
    register_*_ci apply to registered suites under test/registered/, not here.
  • Documentation — not applicable. No user-facing interface, flag or default
    changes; this restores correct output for checkpoints that already load today.
  • Accuracy results above. Speed: unchanged, load-path only.
  • Code style: no duplication, no device synchronization, no in-place argument
    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 checkpoint
tested here has N % 128 == 0, so flatten-only matches. If comfy-kitchen ever
serializes 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 then main gained the
MXFP8OnlineLinearMethod branch (#37903) and the scale_u8 hook argument (#40039);
both are merged into this branch, and the cutlass/cutedsl branch in srt/.../fp8.py
still re-interleaves layer.weight_scale_inv on main, so the bug still reproduces
there.

🤖 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

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>
@TyGu888

TyGu888 commented Sep 4, 2026

Copy link
Copy Markdown
Contributor Author

Per MAINTAINER.md, @mickqian @BBuf are the Diffusion oncalls for python/sglang/multimodal_gen. Could one of you add the run-ci label so CI can run here? No Merge Oncall was auto-assigned on this PR.

The 30-second version: comfy-kitchen serializes MXFP8 block scales already in SWIZZLE_32_4_4 order, so the cutlass/cutedsl branch of Fp8LinearMethod._process_mxfp8_linear_weight_scale runs block_scale_interleave over them a second time. The failure is silent — the checkpoint loads, the kernels run at full speed, and the sample decodes to colour noise. On one real MiniMax-H3 attention matrix, cosine goes 0.739481 → 0.999641 with the fix.

The change is confined to multimodal_gen: comfy-marked MXFP8 checkpoints get a thin Fp8LinearMethod subclass, and config.json-driven MXFP8 keeps SRT's behaviour unchanged (layer_markers is only ever set by resolve_comfy_checkpoint_quantization). The two added unit tests are CPU-only and need no checkpoint fixture, so multimodal-gen-test should cover this without a GPU runner.

Happy to adjust anything — in particular there is an open question in the description about block_scale_interleave padding for shapes where N % 128 != 0.

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>
@mickqian

mickqian commented Oct 6, 2026

Copy link
Copy Markdown
Collaborator

/tag-and-rerun-ci

@github-actions github-actions Bot added the run-ci CI: run the baseline test suite on this PR label Oct 6, 2026
@TyGu888

TyGu888 commented Oct 7, 2026

Copy link
Copy Markdown
Contributor Author

Thanks @mickqian for resolving the conflict with #37903 and carrying the override over to the scale_u8 argument from #40039. The forwarding is what the fix intends: an explicit scale_u8 comes from the block-FP8 conversion in row-major order and still needs SRT's interleave; only the checkpoint's own, already swizzled bytes skip it.

CI on 65f8ce1 has finished. The new tests and the existing MXFP8 tests pass in multimodal-gen-unit-test (and in AMD 1-GPU shard 1), as do the NVIDIA 1-GPU suites (including 5090 and B200) and minimax_h3_ref2va_video_audio_2gpu_h100 on H100. None of the red checks reach the MXFP8 loading path:

  • multimodal-gen-unit-test: the only failure is test_kandinsky6_sr_http_request.py::test_multipart_sr_request_without_prompt_is_accepted, which fails the same way on main at e179358 (reproduced locally). [diffusion] Fix width/height silently dropped on multipart /v1/videos requests #35108 added Form(None) defaults for width/height to create_video, and the test calls it directly, so the Form markers reach VideoGenerationsRequest.
  • multimodal-gen-test-2-gpu (1) was skipped by the fail-fast health check because of that job, which is why diffusion-coverage-check reports 9 missing 2-GPU cases.
  • multimodal-gen-test-2-gpu (2): three LTX-2 two-stage cases miss their H100 perf limits (e.g. average denoise step 312.8 ms vs 300.2 ms).
  • AMD: the 1-GPU shard 0 and 1 failures are ROCm-specific (FA falls back to SDPA, no is_full_nvlink on the ROCm platform, NaN in test_minimax_h3_vae_parallel_modes, exact-equality misses of ~1e-6 in the Kandinsky6 layerwise-offload test) plus the same Kandinsky SR test. In 2-GPU shard 1, minimax_h3_ref2va_video_audio_2gpu_h100 decodes to NaN with the BF16 checkpoint (quantization=None), which looks like the same ROCm VAE NaN, and ltx_2_5_diffusion_decoder_2gpus hits an inductor compile-worker crash.
  • The cancelled XPU job and the two PR Test Extra gates (no run-ci-extra label) are unrelated as well.

@mickqian
mickqian merged commit aa5551d into sgl-project:main Oct 7, 2026
164 of 179 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

diffusion SGLang Diffusion quant LLM Quantization run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants