Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
19 commits
Select commit Hold shift + click to select a range
cb34bdd
perf(diffusion): use custom all-reduce v2 for TP
BBuf Aug 27, 2026
d383209
Fuse Qwen-Image text QKV projections
BBuf Aug 27, 2026
2a37870
Optimize masked diffusion attention packing
BBuf Aug 27, 2026
8e056cf
Merge branch 'main' into perf/diffusion-custom-allreduce-v2
BBuf Aug 27, 2026
e7d6121
Merge branch 'main' into perf/diffusion-custom-allreduce-v2
BBuf Aug 28, 2026
d1e5b4e
Merge branch 'main' into perf/diffusion-custom-allreduce-v2
BBuf Aug 28, 2026
d4edee1
Merge branch 'main' into perf/diffusion-custom-allreduce-v2
BBuf Aug 30, 2026
a5c2e1e
test(diffusion): remove fused QKV CPU test
BBuf Aug 30, 2026
c34959d
test(diffusion): bump Qwen-Image consistency GT revision
BBuf Aug 30, 2026
aa11013
fix(diffusion): preserve Qwen-Image NVFP4 fallback mappings
BBuf Aug 30, 2026
1ed919b
test(diffusion): pin corrected Sana consistency GT
BBuf Aug 31, 2026
b385ffe
Merge main into perf/diffusion-custom-allreduce-v2
BBuf Aug 31, 2026
dfedfc1
Merge branch 'main' into perf/diffusion-custom-allreduce-v2
mickqian Aug 31, 2026
e5b6cdc
Gate Qwen-Image fused added QKV by quality
BBuf Aug 31, 2026
272e6d1
Merge remote-tracking branch 'origin/main' into codex/pr36680-ci-data…
BBuf Aug 31, 2026
990fd1a
Merge origin/main to resolve Qwen-Image TP collective test conflicts.
BBuf Aug 31, 2026
e979cfb
fix(diffusion): restore canonical lossless GT pin
BBuf Aug 31, 2026
a2d110d
Merge origin/main to resolve Qwen-Image TP test conflicts.
BBuf Aug 31, 2026
45da391
Merge origin/main to resolve Qwen-Image TP QKV packing conflicts.
BBuf Sep 1, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 39 additions & 1 deletion docs/cookbook/diffusion/Qwen-Image/Qwen-Image.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -55,7 +55,45 @@ sglang generate \
--save-output
```

### 3.2 Configuration Tips
### 3.2 Fixed-resolution latency on two H200 GPUs

For `Qwen/Qwen-Image-2512` at 1024x1024, use breakable CUDA graph (BCG) to
reduce launch overhead across graph-safe DiT segments while retaining explicit
breakpoints around unsupported operations. This recipe was validated on two
NVIDIA H200 GPUs with 50 denoising steps and no classifier-free guidance:

```bash Command
sglang serve \
--model-path Qwen/Qwen-Image-2512 \
--model-type diffusion \
--num-gpus 2 \
--tp-size 2 \
--performance-mode speed \
--dit-layerwise-offload false \
--enable-torch-compile false \
--enable-breakable-cuda-graph \
--warmup-mode server \
--warmup-resolutions 1024x1024
```

Declare every production resolution in `--warmup-resolutions`. A request at an
uncaptured resolution runs eagerly, so omitting `1024x1024` removes the gain
from this recipe. Graph capture used about 5 GB more peak memory per GPU in the
validation run.

On CUDA, the TP path dispatches supported collectives through SRT
CustomAllReduceV2. At 1024x1024, Qwen-Image reduces 24 MiB row-parallel
outputs; the diffusion runtime reserves a 32 MiB V2 workspace so these
collectives do not fall back to NCCL. If profiling shows large NCCL all-reduce
kernels again, first confirm that V2 is enabled and the requested shape fits
the workspace.

BCG changed floating-point execution order but not the sampling algorithm. The
fixed-seed output measured 0.984 SSIM and 39.7 dB PSNR against eager output; use
eager execution when you require bit-exact output. Regional `torch.compile` was
also tested on this profile and did not improve steady-state latency.

### 3.3 Configuration Tips

Currently supported optimizations are listed [here](/docs/sglang-diffusion/compatibility_matrix).

Expand Down
12 changes: 9 additions & 3 deletions docs/docs/sglang-diffusion/performance-optimization.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,11 @@ These settings should preserve model behavior while changing residency, parallel
<td style={{padding: "9px 12px"}}>You want a safe preset for speed or memory without overriding explicit flags.</td>
<td style={{padding: "9px 12px"}}><a href="./deployment_cookbook">Deployment and Performance Modes</a></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500}}>Breakable CUDA graph</td>
<td style={{padding: "9px 12px"}}>A supported pipeline serves a fixed set of shapes and eager execution is launch-bound.</td>
<td style={{padding: "9px 12px"}}><a href="./api/cli">CLI reference</a></td>
</tr>
<tr>
<td style={{padding: "9px 12px", fontWeight: 500}}>Offload, FSDP, CFG parallelism</td>
<td style={{padding: "9px 12px"}}>GPU memory, multi-GPU residency, or CFG branch splitting is the main bottleneck.</td>
Expand Down Expand Up @@ -119,9 +124,10 @@ These techniques can change the denoising path, numerical representation, or gen

1. Establish a baseline with the target model, resolution, frame count, step count, and GPU type.
2. Select `--performance-mode` and explicit residency or parallelism flags.
3. Tune attention backend and batching for the deployment pattern.
4. Profile if the bottleneck is unclear.
5. Add caching, progressive resolution, or quantization only after comparing output quality against your acceptance target.
3. Compare breakable CUDA graph against eager execution for supported fixed-shape pipelines. Pass every served resolution to `--warmup-resolutions` and confirm capture in the server log.
4. Tune attention backend and batching for the deployment pattern.
5. Profile if the bottleneck is unclear.
6. Add caching, progressive resolution, or quantization only after comparing output quality against your acceptance target.

## Diagnostics

Expand Down
3 changes: 2 additions & 1 deletion python/sglang/kernels/ops/diffusion/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -140,7 +140,8 @@ tensor copy per residual site.
### Data movement (all bit-exact by construction)

`usp_merge_heads`, `pack_qkv_destination_major`, `fused_pack_qkv`,
`fused_scatter_to_padded`, `fused_causal_conv3d_cat_pad_cuda`,
`fused_pack_segmented_qkv`, `fused_scatter_to_padded`,
`fused_causal_conv3d_cat_pad_cuda`,
`cat_pad_channels_last_3d`, `dup_up3d_add`, `fused_temb_table_slices`,
`ltx2_ada_values9`.

Expand Down
12 changes: 12 additions & 0 deletions python/sglang/kernels/ops/diffusion/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -313,6 +313,13 @@
_CUDA,
"Varlen gather of Q/K/V at valid positions.",
),
(
"diffusion.varlen_pack_segmented_qkv",
KernelBackend.TRITON,
"layout.varlen_pack_pad_triton:fused_pack_segmented_qkv",
_CUDA,
"Varlen gather from a virtual prefix/main Q/K/V sequence.",
),
(
"diffusion.varlen_scatter_to_padded",
KernelBackend.TRITON,
Expand Down Expand Up @@ -477,6 +484,7 @@
"usp_merge_heads": "layout.usp_relayout_jit",
"build_inv_indices": "layout.varlen_pack_pad_triton",
"fused_pack_qkv": "layout.varlen_pack_pad_triton",
"fused_pack_segmented_qkv": "layout.varlen_pack_pad_triton",
"fused_scatter_to_padded": "layout.varlen_pack_pad_triton",
"cat_pad_channels_last_3d": "layout.wan_causal_cache_triton",
"dup_up3d_add": "layout.wan_causal_cache_triton",
Expand All @@ -501,6 +509,10 @@
"mount_nvfp4_bias_gelu": "sites.nvfp4_bias_gelu_site",
"nvfp4_bias_gelu_active": "sites.nvfp4_bias_gelu_site",
"unmount_nvfp4_bias_gelu": "sites.nvfp4_bias_gelu_site",
"mark_qwen_image_added_qkv_site": "sites.qwen_image_added_qkv_site",
"mount_qwen_image_added_qkv": "sites.qwen_image_added_qkv_site",
"qwen_image_added_qkv_active": "sites.qwen_image_added_qkv_site",
"unmount_qwen_image_added_qkv": "sites.qwen_image_added_qkv_site",
"can_use_ln_modulate": "sites.fused_ln_modulate_site",
"fused_ln_modulate": "sites.fused_ln_modulate_site",
"fused_ln_modulate_active": "sites.fused_ln_modulate_site",
Expand Down
114 changes: 114 additions & 0 deletions python/sglang/kernels/ops/diffusion/layout/varlen_pack_pad_triton.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,120 @@ def fused_pack_qkv(
)


@triton.jit
def _fused_pack_segmented_qkv_kernel(
Q_prefix_ptr,
K_prefix_ptr,
V_prefix_ptr,
Q_main_ptr,
K_main_ptr,
V_main_ptr,
Q_unpad_ptr,
K_unpad_ptr,
V_unpad_ptr,
indices_ptr,
PREFIX_ROWS,
MAIN_ROWS,
HD,
prefix_row_stride,
main_row_stride,
dst_row_stride,
BLOCK_HD: tl.constexpr,
):
"""Pack a virtual ``[prefix, main]`` sequence without materializing it."""
out_row = tl.program_id(0)
src_row = tl.load(indices_ptr + out_row).to(tl.int64)
joint_rows = PREFIX_ROWS + MAIN_ROWS
batch = src_row // joint_rows
row_in_batch = src_row - batch * joint_rows
from_prefix = row_in_batch < PREFIX_ROWS

prefix_row = batch * PREFIX_ROWS + row_in_batch
main_row = batch * MAIN_ROWS + row_in_batch - PREFIX_ROWS
cols = tl.arange(0, BLOCK_HD)
col_mask = cols < HD
prefix_mask = col_mask & from_prefix
main_mask = col_mask & ~from_prefix

prefix_offset = prefix_row * prefix_row_stride + cols
main_offset = main_row * main_row_stride + cols
dst_offset = out_row * dst_row_stride + cols

q_val = tl.load(Q_prefix_ptr + prefix_offset, mask=prefix_mask, other=0.0)
k_val = tl.load(K_prefix_ptr + prefix_offset, mask=prefix_mask, other=0.0)
v_val = tl.load(V_prefix_ptr + prefix_offset, mask=prefix_mask, other=0.0)
q_val += tl.load(Q_main_ptr + main_offset, mask=main_mask, other=0.0)
k_val += tl.load(K_main_ptr + main_offset, mask=main_mask, other=0.0)
v_val += tl.load(V_main_ptr + main_offset, mask=main_mask, other=0.0)

tl.store(Q_unpad_ptr + dst_offset, q_val, mask=col_mask)
tl.store(K_unpad_ptr + dst_offset, k_val, mask=col_mask)
tl.store(V_unpad_ptr + dst_offset, v_val, mask=col_mask)


def fused_pack_segmented_qkv(
q_prefix: torch.Tensor,
k_prefix: torch.Tensor,
v_prefix: torch.Tensor,
q_main: torch.Tensor,
k_main: torch.Tensor,
v_main: torch.Tensor,
indices: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""Pack Q/K/V from a virtual ``[prefix, main]`` joint sequence.

This is bitwise equivalent to concatenating each prefix/main pair and
calling :func:`fused_pack_qkv`, but skips the three dense concatenations.
All inputs use ``[B, S, H, D]`` layout and share batch/head dimensions.
"""
prefixes = (q_prefix, k_prefix, v_prefix)
mains = (q_main, k_main, v_main)
assert q_prefix.shape == k_prefix.shape == v_prefix.shape
assert q_main.shape == k_main.shape == v_main.shape
assert q_prefix.dim() == q_main.dim() == 4
assert q_prefix.shape[0] == q_main.shape[0]
assert q_prefix.shape[2:] == q_main.shape[2:]
assert all(t.dtype == q_prefix.dtype for t in (*prefixes, *mains))
assert indices.dtype in (torch.int32, torch.int64)

q_prefix, k_prefix, v_prefix = (t.contiguous() for t in prefixes)
q_main, k_main, v_main = (t.contiguous() for t in mains)
prefixes = (q_prefix, k_prefix, v_prefix)
mains = (q_main, k_main, v_main)
batch_size, prefix_rows, num_heads, head_dim = q_prefix.shape
main_rows = q_main.shape[1]
hd = num_heads * head_dim
n_valid = indices.shape[0]
if n_valid == 0:
return tuple(
t.new_empty(0, num_heads, head_dim) for t in (q_prefix, k_prefix, v_prefix)
)

prefix_flat = tuple(t.view(batch_size * prefix_rows, hd) for t in prefixes)
main_flat = tuple(t.view(batch_size * main_rows, hd) for t in mains)
outputs = tuple(
torch.empty(n_valid, hd, dtype=q_prefix.dtype, device=q_prefix.device)
for _ in range(3)
)
block_hd = triton.next_power_of_2(hd)
with torch.get_device_module().device(q_prefix.device):
_fused_pack_segmented_qkv_kernel[(n_valid,)](
*prefix_flat,
*main_flat,
*outputs,
indices,
prefix_rows,
main_rows,
hd,
prefix_flat[0].stride(0),
main_flat[0].stride(0),
outputs[0].stride(0),
BLOCK_HD=block_hd,
)

return tuple(out.view(n_valid, num_heads, head_dim) for out in outputs)


# ---------------------------------------------------------------------------
# Scatter (pad) — write packed output to [B, S, H, D] with zeros at invalid
# ---------------------------------------------------------------------------
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
"""Qwen-Image added-QKV GEMM packing, gated by request quality.

Packing the three BF16 text projections into one GEMM changes the reduction
association and is therefore not bit-exact. The packed weights stay resident
for checkpoint compatibility, but ``quality="lossless"`` applies their three
slices independently. ``quality="high"`` mounts the single-GEMM path.
"""

from __future__ import annotations

import logging

import torch.nn as nn

from sglang.kernels.ops.diffusion.sites.quality_gate import QualityGatedFusion

logger = logging.getLogger(__name__)

_FUSION = QualityGatedFusion(
name="Qwen-Image fused added-QKV",
marker_attr="_sgl_qwen_image_added_qkv_site",
enabled_attr="_sgl_qwen_image_added_qkv_enabled",
)


def mark_qwen_image_added_qkv_site(module: nn.Module) -> None:
"""Mark an unquantized Qwen-Image attention site; it starts unmounted."""
_FUSION.mark(module)


def qwen_image_added_qkv_active(module: nn.Module) -> bool:
"""Whether the request-scoped packed added-QKV GEMM is mounted."""
return _FUSION.is_enabled(module)


def _site_reject_reason(site: nn.Module) -> str | None:
linear = getattr(site, "to_added_qkv", None)
if linear is None:
return "missing to_added_qkv"
if getattr(linear, "quant_config", None) is not None:
return "quantized packed projection"
if len(getattr(linear, "output_partition_sizes", ())) != 3:
return "packed projection does not contain three shards"
return None


def mount_qwen_image_added_qkv(root: nn.Module) -> bool:
return _FUSION.mount(root, reject_reason=_site_reject_reason, logger=logger)


def unmount_qwen_image_added_qkv(root: nn.Module) -> None:
_FUSION.unmount(root)
Original file line number Diff line number Diff line change
Expand Up @@ -378,7 +378,8 @@ Use these as first commands to benchmark, not as universal winners.
| MiniMax-H3 | 1344x768 resolved canvas, 5 seconds / 124 frames at 24 fps, 50 joint video/audio steps | H200: `--num-gpus 4 --ulysses-degree 4 --performance-mode speed --enable-torch-compile false --enable-breakable-cuda-graph false`; H100: TP2 + Ulysses2 | Root ID plus `--model-variant fl2va` for T2VA/FL2VA or `ref2va` for Ref2VA. Ulysses only; no Ring/CFG/SageAttention. Preserve tiled video-VAE decode. BCG is not part of the validated H3 recipe: warmup and serving can have different packed host boundaries, and a replay-capable experiment must still beat eager without excessive graph memory. Profile joint denoise, video VAE, audio VAE/vocoder, encoder, and collectives separately. |
| FLUX.1 / FLUX.2 image | 1024x1024, runtime-default steps/guidance, 1 GPU | `--enable-torch-compile --warmup-mode request --dit-layerwise-offload false` | `black-forest-labs/FLUX.*` repos are gated; for FP8/NVFP4 use validated `--transformer-path` or `--transformer-weights-path` flows from the quant skill. |
| FLUX.2 Klein / Klein Base | 1024x1024, runtime-default steps/guidance, 1 GPU | `--enable-torch-compile --warmup-mode request --dit-layerwise-offload false` | Current registry has `black-forest-labs/FLUX.2-klein-4B`, `FLUX.2-klein-9B`, and base variants. Klein is step-distilled; Klein Base is not. |
| Qwen-Image / Qwen-Image-Edit | 1024x1024, runtime-default steps/guidance, 1 GPU | `--enable-torch-compile --warmup-mode request`; optionally native `SGLANG_CACHE_DIT_ENABLED=true` | Cache-DiT is lossy. For edit tasks, keep reference image, seed, and output size fixed. |
| Qwen-Image / Qwen-Image-2512 | 1024x1024, 50 steps, no CFG, 2x H200 | `--num-gpus 2 --tp-size 2 --performance-mode speed --dit-layerwise-offload false --enable-torch-compile false --enable-breakable-cuda-graph --warmup-mode server --warmup-resolutions 1024x1024` | Validated on H200. BCG reduced median denoise time from 124.7 to 83.1 ms/step in the same-topology run. Capture every served resolution; an uncaptured shape runs eagerly. CUDA TP should select CustomAllReduceV2 with a 32 MiB diffusion workspace: the 1024x1024 row-parallel outputs are 24 MiB and otherwise fall back to NCCL. Capture used about 5 GB more peak memory per GPU. Fixed-seed output versus eager measured 0.984 SSIM / 39.7 dB PSNR but was not bit-exact. Establish an eager baseline and remeasure BCG on other hardware or shapes. Cache-DiT remains lossy. |
| Qwen-Image-Edit | 1024x1024, runtime-default steps/guidance, 1 GPU | Start eager, then compare `--enable-torch-compile --warmup-mode request` | Keep the reference image, seed, and output size fixed. Do not transfer the Qwen-Image-2512 BCG result without a model-backed edit test. |
| Krea-2 | 1024x1024, distilled `oss_turbo` defaults (8 steps, guidance 1.0) | `--performance-mode speed --warmup-mode request` | Native `krea/Krea-2` text-to-image path with Qwen3-VL text conditioning. The repo may require HF access; keep the 8-step distilled baseline separate from non-turbo sampling experiments. |
| Z-Image / Z-Image-Turbo | 1024x1024, runtime-default steps/guidance, 1 GPU | `--enable-torch-compile --warmup-mode request` | Keep base Z-Image separate from Turbo: base uses 50-step CFG defaults, Turbo uses 9-step zero-CFG defaults. Mainline has bf16-native Triton RMSNorm scale and tanh-residual fusions. |
| Wan2.2 A14B T2V/I2V | 1280x720, 81 frames | Nightly: `--num-gpus 4 --enable-cfg-parallel --ulysses-degree 2 --text-encoder-cpu-offload --pin-cpu-memory` | For lowest latency, also benchmark pure Ulysses on the same GPUs. |
Expand Down
24 changes: 24 additions & 0 deletions python/sglang/multimodal_gen/configs/models/dits/qwenimage.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,25 @@ class QwenImageArchConfig(DiTArchConfig):

param_names_mapping: dict = field(
default_factory=lambda: {
# Merge the short text-stream projections into one tensor-parallel
# GEMM. The loader only applies these rules when the fused target
# exists, so quantization backends that keep the original modules
# continue to load their unfused parameters.
r"^(.*\.attn)\.add_q_proj\.(.+)$": (
r"\1.to_added_qkv.\2",
0,
3,
),
r"^(.*\.attn)\.add_k_proj\.(.+)$": (
r"\1.to_added_qkv.\2",
1,
3,
),
r"^(.*\.attn)\.add_v_proj\.(.+)$": (
r"\1.to_added_qkv.\2",
2,
3,
),
# LoRA mappings
r"^(transformer_blocks\.\d+\.attn\..*\.lora_[AB])\.default$": r"\1",
# SVDquant mappings
Expand All @@ -35,6 +54,11 @@ class QwenImageArchConfig(DiTArchConfig):
}
)

# Serialized ModelOpt checkpoints keep the added Q/K/V projections as
# separate modules, including their BF16 fallback layers. Do not apply the
# runtime-only fused mapping while inferring their quantized tensor layout.
quant_param_names_mapping: dict = field(default_factory=dict)

def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
Expand Down
Loading
Loading