Repository navigation
[diffusion] model: support VDN-H3 (hybrid window softmax + Video Delta linear attention MiniMax-H3, 8-NFE distill) with a hybrid_window_attn_h3 backend - #37903
Merged
Conversation
kevin-mii
requested review from
AgainstEntropy,
BBuf,
DarkSharpness,
HaiShaw,
HydraQYH,
JustinTong0323,
OrangeRedeng,
celve,
mickqian,
ping1jing2,
sogalin,
wisclmy0611,
yichiche,
yingluosanqian,
yuan-luo and
zijiexia
as code owners
September 4, 2026 02:34
Collaborator
|
/tag-and-rerun-ci |
2 tasks done
…dule Add VDNHybridAttentionArchConfig (chunk/radius window, anchors, delta rule, short conv, text state) populated from transformer/config.json's hybrid_attention key, name-mapping rules for the VDN linear-branch keys, and MiniMaxH3VDNLinearBranch: the eager port of OpenVDN's bidirectional Video Delta rule branch (features, frame statistics, Cholesky delta rule, forward/reverse scans, alpha-bridged boundary gather, gated RMSNorm readout). Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01M1KA1GB6xFrCp2qbMy9u23
…ion, pipeline wiring - HybridWindowAttentionH3Backend: chunk-aligned window softmax over the packed H3 layout as a union of dense FlashAttention varlen calls (dense-query rows for text/audio/anchor frames, per-chunk gathered K/V for the rest), gate epilogue, request-static metadata (row-group plan, layout, full-seq RoPE cache for Ulysses); enum + CUDA resolver (SM90/SM100/SM103). - MiniMaxH3Attention: softmax_gate, linear_attention, to_out_linear for DiT blocks when the arch carries hybrid_attention; a dedicated eager hybrid core keeps raw q/k for the branch, applies QK-norm+RoPE out of place, exchanges beta/gates by head and reduces the frame mean under Ulysses. - Denoising stage builds the hybrid metadata once per request. - VDNH3Pipeline / VDNH3PipelineConfig (forces the hybrid backend, rejects --model-variant, ring, compile, BCG) / VDNH3SamplingParams (9 sigma grid points = 8 NFE, t2va only); registry + overlay registry entries. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01M1KA1GB6xFrCp2qbMy9u23
Window backend vs masked-dense reference on a ragged layout (both kernel paths), full-cover == dense, gate epilogue, refiner dense fallback; linear branch arithmetic vs step-by-step references (statistics, Cholesky delta rule, scans, exact-complement gather, text-state decay, head-slice == full run, module vs from-scratch algorithm); registry/admission; the overlay prefuse on synthetic tensors. vdn-h3 / vdn-h3-fp8 bench presets. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01M1KA1GB6xFrCp2qbMy9u23
…tests until the kernel lands Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01M1KA1GB6xFrCp2qbMy9u23
… auxiliary components fall back Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01M1KA1GB6xFrCp2qbMy9u23
…rs on the hybrid backend Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01M1KA1GB6xFrCp2qbMy9u23
BBuf
approved these changes
Sep 9, 2026
4 tasks done
The AMD multimodal-gen jobs failed on eight of the new tests. All eight are ROCm-only: `hybrid_window_attn_h3` is registered by the CUDA platform alone (`RocmPlatform.get_attn_backend_cls_str` rejects it, and `RocmPlatform` has no `_prepare_flash_attention_for_blackwell`), but the guards used `torch.cuda.is_available()`, which is also true under ROCm. - test_hybrid_window_h3_attention.py: gate `requires_cuda` on `current_platform.is_cuda()` instead. - test_minimax_h3_vdn.py: add `requires_cuda_backend` to the two admission tests that resolve the backend through `get_attn_backend`. - test_mxfp8_online_gemm.py: the unaligned-fallback layer used K=8; the fallback's scaled GEMM still wants K % 16 == 0, which ROCm enforces. K=16 is still not a multiple of 32, so the layer keeps the per-channel path and the test keeps its meaning on both backends. Verified on B200 (sm100): 32 passed, 1 skipped across the three files, the same counts as before the change, so no NVIDIA coverage is lost. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Collaborator
Author
|
/rerun-failed-ci |
1 similar comment
Collaborator
Author
|
/rerun-failed-ci |
This was referenced Sep 12, 2026
4 of 5 tasks
Rockdu
added a commit
to Rockdu/sglang
that referenced
this pull request
Sep 23, 2026
… path sgl-project#37903 added _accepts_mxfp8_input, which reads linear.quant_method. A LoRA wrapper (RowParallelLinearWithLoRA) has no such attribute and must run its own forward to apply the delta, so it now reports False. This broke miles_diffusion test_h3_t2va_grpo_2xGPU on sglang-miles. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Rockdu
added a commit
to Rockdu/sglang
that referenced
this pull request
Sep 24, 2026
… path sgl-project#37903 added _accepts_mxfp8_input, which reads linear.quant_method. A LoRA wrapper (RowParallelLinearWithLoRA) has no such attribute and must run its own forward to apply the delta, so it now reports False. This broke miles_diffusion test_h3_t2va_grpo_2xGPU on sglang-miles. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Rockdu
added a commit
to Rockdu/sglang
that referenced
this pull request
Sep 24, 2026
… path sgl-project#37903 added _accepts_mxfp8_input, which reads linear.quant_method. A LoRA wrapper (RowParallelLinearWithLoRA) has no such attribute and must run its own forward to apply the delta, so it now reports False. This broke miles_diffusion test_h3_t2va_grpo_2xGPU on sglang-miles. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
1 task done
This was referenced Oct 4, 2026
mickqian
added a commit
to TyGu888/sglang
that referenced
this pull request
Oct 6, 2026
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>
4 of 5 tasks
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Motivation
Add OpenVDN/vdn-minimax-h3 (VDN-H3): MiniMax-H3 with hybrid attention (chunked window softmax + a Video Delta linear branch), 8-NFE DMD2 distill, t2va only, with the
hybrid_window_attn_h3backend it was trained for.8x B200, 1344x768 at 24 fps with audio, 14.375 s (345 frames, ~104k packed rows),
num_inference_steps: 9(8 DiT forwards), seed 1000, Ulysses 8, eager, warmup at the served clip shape:mxfp8)Modifications
VDNH3Pipeline/VDNH3PipelineConfig/VDNH3SamplingParamslive in their ownminimax_h3_vdn*modules next to the H3 ones. The base H3 code changes in four places: the arch config gains ahybrid_attentionfield and one name-mapping rule,MiniMaxH3Attentionbuilds onehybridsubmodule and dispatches to it when the hybrid backend is selected,MiniMaxH3MLPhandsfc2the fused SwiGLU+quant output when its quant method accepts it, and the denoising stage asks the VDN module for the request-static attention metadata. A revision-pinned model overlay (kevin-mi/VDN-H3-overlay@2bf8c736) materializes the repo into the base-H3 layout: prefuseddefault+turboLoRAs, the linear-branch weights,hybrid_attentionin the transformer config. Admission forceshybrid_window_attn_h3for the transformer (a dense backend would silently skip the linear branch and the gates) and rejects--model-variant,quality="high", ring,torch.compileand breakable CUDA graph.hybrid_window_attn_h3. Chunk-5 / radius-1 window with anchor frames as a union of FlashAttention varlen calls (dense-query rows against all keys, per-chunk gathered K/V passes bounded byvdn_max_gather_rows), softmax gate epilogue, request-static metadata built once per request.minimax_h3_vdn.py): short conv, TF32 frame statistics, Cholesky delta rule, forward/reverse scans, alpha-bridged boundary gather, gated RMSNorm readout, text-state seed; per-head independent, so Ulysses shards it exactly. Fused Triton kernels for the conv + SiLU + L2 norm (q written frame-major for the readout), statistics prologue, boundary gather and epilogue. The scans fold each chunk into one affine map and run one chain over the chunks with both directions per launch, since the chunked gather reads states at chunk boundaries only; the prompt's text state joins the frames' Cholesky/solve batch as a leading virtual frame.minimax_h3_vdn_attention.py): out-of-place fused QK-norm + RoPE (a JIT variant of the H3 kernel, bit-equal to the in-place one) so the branch keeps the raw q/k. Under Ulysses q, k, v and the per-head scalars travel as four async field-major all-to-alls that land contiguous; only the 128-wide output-gate hidden is all-gathered (upruns on the head shard); the frame mean is a reshape-sum all-reduce. Served-shape warmup (--warmup-num-frames/--warmup-resolutions) landed separately in [diffusion] MiniMax-H3: warm up at the served clip shape #37945.MXFP8OnlineLinearMethod(mxfp8_online.py) sits behind the existingMXFP8Configwhen the checkpoint is not fp8-serialized. Weights quantize at load to e4m3 with one E8M0 scale per 32 elements in the cuBLASLt swizzled layout, and the GEMM is cuBLASLt's block-scaledtorch.nn.functional.scaled_mm. The activation quant is fused into the producers (Triton kernels inkernels/ops/diffusion/quantization/, each byte-exact againstflashinfer.mxfp8_quantizeof the bf16 result): adaLN modulation + quant for the qkv and fc1 inputs (the qkv site keeps the bf16 rows for the VDN branch), SwiGLU + quant for fc2, and a plain quantizer for out_proj; linears take the(fp8, scales)tuple throughLinearMethodBase.accepts_mxfp8_input. Layers with K not a multiple of 32 or fp32 params keep the per-channel path. VDN-H3 defaults to it on SM100+ (--quantization fp8selects it too,--quantization bf16opts out; before SM100fp8stays per-channel); other models are unchanged.test_minimax_h3_vdn.py(branch algorithm vs a from-scratch reference, boundary scans and fused kernels vs the plain chain through the real forward, head-slice contract, gather invariants, frame sums, config/registry contracts),test_hybrid_window_h3_attention.py,test_vdn_ulysses_exchange_2_gpu.py,test_mxfp8_online_gemm.py, the registered kernel teststest_vdn_linear_branch.py,test_mxfp8_swizzled.pyandtest_qknorm_rope_out_of_place.py(each kernel has acan_use_*predicate and a README row), GPU casevdn_h3_t2va_4gpu_h100, cookbook section and B200 / RTX PRO 6000 tables, attention-backend entry, thevdn-h3benchmark preset.Accuracy Tests
Block-level parity against OpenVDN's
HybridAttentionon blocks 0 / 25 / 49 with the prefused weights: relative L2 5e-3 to 7e-3 (the same band as the dense base+LoRA smoke). Ulysses 2 and 4 against the single-process full-sequence path on one real block: relative L2 2e-5, bitwise on most ranks. MXFP8, per-tensor and per-channel fp8 all sit at 5.8% relative L2 from bf16 on that block. The branch with the boundary scans and fused kernels matches the plain frame-chain / eager chain through the real forward, including the anchor-frame shift of the chunk grid.Note for reviewers: two identical 8-GPU runs of the 8-step sampler agree at only ~19 dB PSNR (bf16 and atomic reduction order amplified over 8 steps), so end-to-end PSNR is not a meaningful gate for this model.
Speed Tests and Profiling
Profile-driven on 8x B200 (
--profile --num-profiled-timesteps 2, kernels attributed to aten ops): the first port ran 1.24 s/NFE with ~30k launches per forward, half of them 6 us fp32 GEMMs from the 102-step frame scan. Chunk-boundary scans, contiguous K/V gathers, gate hidden + reshape-sum frame means, fused SwiGLU/quant + batched chains, the field-major exchange, MXFP8 into cuBLASLt with the quant fused into the producers, and the text state batched into the frame solve bring it to 0.93 s/NFE. Rank profile before the last two: 32% window FlashAttention, 24% fp8 GEMM, 12% NCCL, ~19.7k launches per forward.Measured and not shipped: OpenVDN's branch-parallel layout (softmax and linear branches on disjoint rank groups) at 1.39 s/NFE (5+3) and 1.49 (6+2); two-group head-pipelined all-to-alls at 1.10; a static-tile Triton window kernel at 226 vs 169 ms per block; breakable CUDA graph capture runs out of memory at 104k rows. MXFP8 GEMM backends at the fc1/qkv shapes: cuBLASLt 3.40 ms, FlashInfer cuDNN 3.61, DeepGEMM 3.73, FlashInfer CUTLASS 4.18, CuTe-DSL 5.06, TRT-LLM 17.9 (decode-shaped).
Deliberately not in this PR
Known limitations
stage-b-step-2000) is out of scope.Checklist
🤖 Generated with Claude Code
CI States
Latest PR Test (Base): ⏳ Run #34594729897
Latest PR Test (Extra): ❌ Run #34594729445
Latest PR Test (AMD ROCm 10): ❌ Run #34594729539