Skip to content

[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
mickqian merged 77 commits into
sgl-project:mainfrom
kevin-mii:VDN-h3
Sep 12, 2026

Conversation

@kevin-mii

@kevin-mii kevin-mii commented Sep 4, 2026 •

Copy link
Copy Markdown
Collaborator

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_h3 backend 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:

Config s/NFE Denoise (8 forwards) Request after warmup Peak HBM
OpenVDN published (fp8, 5+3 branch-parallel Ulysses) 1.40 11.2 s - -
SGLang default (online mxfp8) 0.93 7.45 s 9.3 s 77.7 GB/GPU

Modifications

  • Model. VDNH3Pipeline / VDNH3PipelineConfig / VDNH3SamplingParams live in their own minimax_h3_vdn* modules next to the H3 ones. The base H3 code changes in four places: the arch config gains a hybrid_attention field and one name-mapping rule, MiniMaxH3Attention builds one hybrid submodule and dispatches to it when the hybrid backend is selected, MiniMaxH3MLP hands fc2 the 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: prefused default + turbo LoRAs, the linear-branch weights, hybrid_attention in the transformer config. Admission forces hybrid_window_attn_h3 for the transformer (a dense backend would silently skip the linear branch and the gates) and rejects --model-variant, quality="high", ring, torch.compile and 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 by vdn_max_gather_rows), softmax gate epilogue, request-static metadata built once per request.
  • Linear branch (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.
  • Hybrid attention module (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 (up runs 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.
  • fp8. Online MXFP8 for bf16 checkpoints: MXFP8OnlineLinearMethod (mxfp8_online.py) sits behind the existing MXFP8Config when 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-scaled torch.nn.functional.scaled_mm. The activation quant is fused into the producers (Triton kernels in kernels/ops/diffusion/quantization/, each byte-exact against flashinfer.mxfp8_quantize of 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 through LinearMethodBase.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 fp8 selects it too, --quantization bf16 opts out; before SM100 fp8 stays per-channel); other models are unchanged.
  • Tests and docs. 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 tests test_vdn_linear_branch.py, test_mxfp8_swizzled.py and test_qknorm_rope_out_of_place.py (each kernel has a can_use_* predicate and a README row), GPU case vdn_h3_t2va_4gpu_h100, cookbook section and B200 / RTX PRO 6000 tables, attention-backend entry, the vdn-h3 benchmark preset.

Accuracy Tests

Block-level parity against OpenVDN's HybridAttention on 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

  • A switch for the branch-overlap schedule under Ulysses: the alternatives were measured (8x B200, steady s/NFE: no overlap 0.890, reordered collective 0.879, side stream 0.855; output bit-identical) and only the side-stream schedule ships.

Known limitations

  • t2va only; fl2va / ref2va were not trained. TP > 1 untested for the branch parameters.
  • The 50-step checkpoint (stage-b-step-2000) is out of scope.
  • License note for maintainers: the checkpoint inherits the MiniMax-H3 Community License.

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

@mickqian

mickqian commented Sep 4, 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 Sep 4, 2026
kevin-mii and others added 6 commits September 4, 2026 00:49
…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
kevin-mii and others added 5 commits September 10, 2026 16:59
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>
@kevin-mii

Copy link
Copy Markdown
Collaborator Author

/rerun-failed-ci

1 similar comment
@kevin-mii

Copy link
Copy Markdown
Collaborator Author

/rerun-failed-ci

@mickqian
mickqian merged commit ff1ce11 into sgl-project:main Sep 12, 2026
452 of 532 checks passed
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>
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

diffusion SGLang Diffusion documentation Improvements or additions to documentation jit-kernel 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.

4 participants