Skip to content

[Diffusion] Keep the Wan VAE decoder channels_last and add a Triton NHWC nearest upsample - #38182

Merged
BBuf merged 4 commits into
sgl-project:mainfrom
Dayuxiaoshui:qwen-image-vae-fast-path
Sep 9, 2026
Merged

BBuf merged 4 commits into
sgl-project:mainfrom
Dayuxiaoshui:qwen-image-vae-fast-path

Conversation

@Dayuxiaoshui

@Dayuxiaoshui Dayuxiaoshui commented Sep 6, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

Follow-up to #38020. The degenerate-stride layout issue found on the Qwen-Image VAE also affects the Wan 2.1 video VAE, where the decoder is a much larger share of a request: on TurboWan2.1-T2V-1.3B (4 steps, 81 frames, 480x832) a warm request spends ~0.93 s denoising and ~1.56 s in DecodingStage, i.e. the VAE is ~60% of the request.

Two things were going on in the Wan decoder at quality=high:

  1. resample_forward's upsample3d frame interleave (reshape / stack / reshape) materialises NCDHW, and the attention residual's operand order drops channels_last_3d too, so every up block started with an NCDHW tensor: 4 of the 29 RMSNorm+SiLU sites fell back to eager per chunk, the residual adds ran strided, and cuDNN converted the resample conv2d input to NHWC and back.
  2. Once the decoder does run channels_last, aten's upsample_nearest2d_nhwc kernel is several times slower than its NCHW sibling on the same bytes (0.458 ms vs 0.190 ms for [1, 192, 240, 416] bf16 on H200); the decoder hits it once per up block per chunk.

Modifications

New kernel, lossless (bit-exact, unconditional)

  • kernels/ops/diffusion/layout/nearest_upsample_nhwc_triton.py: integer-factor nearest upsample of a dense channels_last [N, C, H, W] tensor as a pure gather. For an integer factor f, writing i = k*f + r with 0 <= r <= f-1, nearest reads floor(i/f) = k and nearest-exact reads floor((i+0.5)/f) = k since (r+0.5)/f < 1, so both read i // f and the result is bitwise identical to nn.Upsample in either mode. Walks the output in NHWC memory order so loads and stores are contiguous along C. 0.061 ms for the shape above. Registered as diffusion.nearest_upsample_nhwc with can_use_nearest_upsample_nhwc; raises on unsupported input.
  • GatedChannelsLastUpsample runs it whenever the upsample input already has canonical NHWC strides. In that case aten itself would run its NHWC kernel and return the same dense channels_last tensor, so the replacement is layout- and value-identical and is enabled on the lossless path too.

Quality-gated (extra-high / high)

  • wanvae.py: _interleave_time_pairs writes the interleaved frames straight into a channels_last_3d buffer with one copy_ (both views are valid: the channel dim has stride 1 so [B, 2C, ...] -> [B, 2, C, ...] is a view; the output's 2T dim splits into (T, 2)), replacing the NCDHW-materialising stack. Gated because the NHWC conv2d it enables need not pick the same cuDNN algorithm.
  • The Wan installer wraps WanUpsample with GatedChannelsLastUpsample (stride canonicalisation, from [Diffusion] Port the Wan VAE decoder fast paths to the Qwen-Image VAE #38020) and hands the gate to WanResample and WanAttentionBlock (residual operand order). All 29 norm sites now hit the fused kernel. The operand order is commutative but changes the reduction order of the eager norm that follows under fp32 autocast, which is how Wan pipelines decode (fp32 VAE weights, bf16 autocast); measured end-to-end as non-bit-exact, hence gated.

Tests / docs

  • test_layout.py: nearest_upsample_nhwc vs F.interpolate for bf16/fp16/fp32, factors 2 and (3, 2), nearest and nearest-exact, N = 1/2/4; predicate rejections.
  • test_model_fast_paths.py: interleave vs stack (values and layout), upsample wrapper dispatch on canonical / degenerate / NCHW inputs with gate on and off, installer wiring.
  • README selection matrix and existing-fast-paths.md updated.

Accuracy Tests

1x H200, Wan 2.1 VAE (Wan-AI/Wan2.1-I2V-14B-480P-Diffusers VAE weights), 81 frames 480x832.

Lossless path

  • Standalone decode, bf16 weights: torch.equal vs main output: True.
  • Standalone decode, fp32 weights + bf16 autocast (the Wan T2V pipeline setting, vae_precision=fp32, vae_decode_precision=bf16): torch.equal vs main: True.
  • End-to-end sglang generate on TurboWan2.1-T2V-1.3B, 6 prompts, seed 1, 4 steps: all 6 mp4s md5-identical to main.
  • Peak decode memory unchanged: 4.577 GB in both modes.

Quality path

  • main-high vs this-PR-high (DiT path identical and confirmed deterministic across runs), 81 frames: PSNR mean 36.3 dB, min 35.0 dB. The extra difference vs main-high comes from the 4 additional norm sites per chunk now reaching the pre-existing wan_rmsnorm_silu kernel; the layout-only part was measured at 89.6 dB on Qwen-Image in [Diffusion] Port the Wan VAE decoder fast paths to the Qwen-Image VAE #38020.

Speed Tests and Profiling

Standalone VAE decode (torch.profiler GPU kernel self time, 81f 480x832):

weights / autocast path main this PR
bf16 lossless 1342 ms / 7850 kernels 1324 ms / 7850
bf16 high 1178 ms / 4580 1112 ms / 4006
fp32 + bf16 autocast (pipeline setting) lossless 1633 ms 1593 ms
fp32 + bf16 autocast high 1311 ms 1215 ms

End-to-end, TurboWan2.1-T2V-1.3B, 4 steps, 81f 480x832, idle H200, main and this PR run interleaved (main from a git worktree of the base commit), 6 prompts per run, median of the 5 warm requests:

stage main this PR
DecodingStage, lossless 1565.7 ms 1531.5 ms (-2.2%)
DecodingStage, high 1061.7 ms 942.7 ms (-11.2%)
DmdDenoisingStage, lossless 925.5 ms 929.4 ms
DmdDenoisingStage, high 883.9 ms 887.4 ms

Denoising is unchanged, so the effect is confined to the VAE. With decode at ~60% of a warm request, high saves ~5% of the request and lossless ~1.4%.

Qwen-Image also picks up the Triton upsample on its high path: 44.6 -> 42.9 ms at 1024x1024.

Checklist

Review and Merge Process

  1. Ping Merge Oncalls to start the process. See the PR Merge Process.
  2. Get approvals from CODEOWNERS and other reviewers.
  3. Trigger CI tests with comments or contact authorized users to do so.
    • Common commands include /tag-and-rerun-ci, /tag-run-ci-label, /rerun-failed-ci
  4. After green CI and required approvals, ask Merge Oncalls or people with Write permission to merge the PR.

CI States

Latest PR Test (Base): 🚫 Run #34174270196
Latest PR Test (Extra): ❌ Run #34174270292
Latest PR Test (AMD ROCm 7.2): ❌ Run #34174270249

…HWC nearest upsample

Follow-up to sgl-project#38020. The degenerate-stride layout issue found on the
Qwen-Image VAE also affects the Wan 2.1 video VAE, where the decoder is a
much larger share of short-step video pipelines, and it comes with a second
problem: once the decoder does run channels_last, aten's
upsample_nearest2d_nhwc kernel is several times slower than its NCHW
sibling on the same bytes.

Lossless (bit-exact, unconditional):
- New kernel layout/nearest_upsample_nhwc_triton.py: integer-factor nearest
  upsample of a dense channels_last [N, C, H, W] tensor as a pure gather.
  For integer factors nearest and nearest-exact both map output index i to
  input index i // f, so the result is bitwise identical to nn.Upsample in
  either mode for every dtype. Registered as diffusion.nearest_upsample_nhwc
  with a can_use_ predicate. 0.061 ms vs 0.458 ms (aten NHWC) and 0.190 ms
  (aten NCHW) for [1, 192, 240, 416] bf16 on H200.
- GatedChannelsLastUpsample runs it whenever the upsample input already has
  canonical NHWC strides (aten would run its NHWC kernel and return the
  same dense channels_last tensor), on the lossless path too.

Quality-gated (extra-high / high):
- resample_forward's upsample3d frame interleave (reshape / stack / reshape,
  which materialises NCDHW) writes the same values straight into a
  channels_last_3d buffer with one copy via _interleave_time_pairs.
- The Wan installer now wraps WanUpsample with GatedChannelsLastUpsample and
  hands the gate to WanResample and WanAttentionBlock (residual operand
  order), so all 29 RMSNorm+SiLU sites hit the fused kernel instead of 25.
  The attention operand order is commutative but changes the reduction order
  of the eager norm that follows under fp32 autocast, which is how Wan
  pipelines decode (fp32 VAE weights, bf16 autocast); measured end-to-end as
  non-bit-exact, hence gated.

Wan 2.1 VAE, 81 frames 480x832 on H200, torch.profiler GPU kernel time:
bf16 weights: main lossless 1342 ms / high 1178 ms -> 1324 / 1112 ms.
fp32 weights + bf16 autocast (the T2V pipeline setting): main 1633 / 1311 ms
-> 1593 / 1215 ms. Lossless mp4 from TurboWan2.1-T2V-1.3B md5-identical to
main; high vs main-high PSNR 36.3 dB mean / 35.0 dB min over 81 frames.
Peak memory unchanged (4.577 GB). Qwen-Image high: 44.6 -> 42.9 ms.
@github-actions github-actions Bot added documentation Improvements or additions to documentation diffusion SGLang Diffusion jit-kernel labels Sep 6, 2026
Comment thread python/sglang/kernels/ops/diffusion/layout/nearest_upsample_nhwc_triton.py Outdated
BBuf and others added 3 commits September 7, 2026 16:59
…contract

Review follow-ups:
- can_use_nearest_upsample_nhwc never raises: _integer_scale rejects
  non-numeric entries, NaN and +/-Inf (and bools) instead of throwing.
- The predicate now requires C > 1 and canonical NHWC strides on every dim.
  That is exactly the condition under which aten's suggest_memory_format
  picks channels_last, so the Triton result is layout-identical as well as
  value-identical; with C == 1 the tensor is also NCHW-contiguous and aten
  returns e.g. strides (4, 4, 2, 1) for a [1, 1, 2, 2] output.
- nearest_upsample_nhwc validates before launching and raises ValueError
  for requires_grad inputs (a direct call must not return a detached
  tensor), unsupported dtype/device/rank, and non-canonical layouts.
- GatedChannelsLastUpsample relies on the predicate alone for the lossless
  dispatch instead of duplicating the stride check.
- Tests assert out.stride() == ref.stride(), cover the malformed scale
  factors, the C == 1 case and the direct-call error paths.
@mickqian

mickqian commented Sep 8, 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 8, 2026
@BBuf
BBuf merged commit db1de6f into sgl-project:main Sep 9, 2026
356 of 427 checks passed
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 run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants