diff --git a/docs/cookbook/diffusion/Qwen-Image/Qwen-Image.mdx b/docs/cookbook/diffusion/Qwen-Image/Qwen-Image.mdx
index 4703ca5e1..b73d47a21 100644
--- a/docs/cookbook/diffusion/Qwen-Image/Qwen-Image.mdx
+++ b/docs/cookbook/diffusion/Qwen-Image/Qwen-Image.mdx
@@ -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).
diff --git a/docs/docs/sglang-diffusion/performance-optimization.mdx b/docs/docs/sglang-diffusion/performance-optimization.mdx
index e0206083d..049936e88 100644
--- a/docs/docs/sglang-diffusion/performance-optimization.mdx
+++ b/docs/docs/sglang-diffusion/performance-optimization.mdx
@@ -41,6 +41,11 @@ These settings should preserve model behavior while changing residency, parallel
You want a safe preset for speed or memory without overriding explicit flags. |
Deployment and Performance Modes |
+
+ | Breakable CUDA graph |
+ A supported pipeline serves a fixed set of shapes and eager execution is launch-bound. |
+ CLI reference |
+
| Offload, FSDP, CFG parallelism |
GPU memory, multi-GPU residency, or CFG branch splitting is the main bottleneck. |
@@ -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
diff --git a/python/sglang/kernels/ops/diffusion/README.md b/python/sglang/kernels/ops/diffusion/README.md
index 21d74b714..9653a18a0 100644
--- a/python/sglang/kernels/ops/diffusion/README.md
+++ b/python/sglang/kernels/ops/diffusion/README.md
@@ -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`.
diff --git a/python/sglang/kernels/ops/diffusion/__init__.py b/python/sglang/kernels/ops/diffusion/__init__.py
index fb55b7221..ba44c8aae 100644
--- a/python/sglang/kernels/ops/diffusion/__init__.py
+++ b/python/sglang/kernels/ops/diffusion/__init__.py
@@ -313,6 +313,13 @@ _SPECS: tuple[tuple[str, KernelBackend, str, frozenset, str], ...] = (
_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,
@@ -477,6 +484,7 @@ _EXPORTS: dict[str, str] = {
"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",
@@ -501,6 +509,10 @@ _EXPORTS: dict[str, str] = {
"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",
diff --git a/python/sglang/kernels/ops/diffusion/layout/varlen_pack_pad_triton.py b/python/sglang/kernels/ops/diffusion/layout/varlen_pack_pad_triton.py
index 937c3ad03..9a1a7f5ad 100644
--- a/python/sglang/kernels/ops/diffusion/layout/varlen_pack_pad_triton.py
+++ b/python/sglang/kernels/ops/diffusion/layout/varlen_pack_pad_triton.py
@@ -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
# ---------------------------------------------------------------------------
diff --git a/python/sglang/kernels/ops/diffusion/sites/qwen_image_added_qkv_site.py b/python/sglang/kernels/ops/diffusion/sites/qwen_image_added_qkv_site.py
new file mode 100644
index 000000000..ab8234f33
--- /dev/null
+++ b/python/sglang/kernels/ops/diffusion/sites/qwen_image_added_qkv_site.py
@@ -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)
diff --git a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-performance/SKILL.md b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-performance/SKILL.md
index d922fc792..4cd1d86ab 100644
--- a/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-performance/SKILL.md
+++ b/python/sglang/multimodal_gen/.claude/skills/sglang-diffusion-performance/SKILL.md
@@ -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. |
diff --git a/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py b/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py
index cb1b15b39..14375c9dd 100644
--- a/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py
+++ b/python/sglang/multimodal_gen/configs/models/dits/qwenimage.py
@@ -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
@@ -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
diff --git a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py
index 1b660f9bd..9e06b97fb 100644
--- a/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py
+++ b/python/sglang/multimodal_gen/runtime/distributed/group_coordinator.py
@@ -42,6 +42,13 @@ logger = init_logger(__name__)
TensorMetadata = namedtuple("TensorMetadata", ["device", "dtype", "size"])
+# Diffusion image tokens make the output of a TP row-parallel projection much
+# larger than the token batches typically seen by the SRT custom all-reduce.
+# Qwen-Image at 1024x1024, for example, reduces 24 MiB tensors. Keep those on
+# the tuned CUDA kernel instead of falling back to NCCL at the default 16 MiB
+# workspace limit.
+_DIFFUSION_CUSTOM_AR_MAX_SIZE = 32 * 1024 * 1024
+
_group_name_counter: dict[str, int] = {}
@@ -239,14 +246,33 @@ class GroupCoordinator:
self.use_custom_op_call = False
def _init_srt_custom_allreduce(self) -> None:
- from sglang.srt.distributed.device_communicators.custom_all_reduce import (
- CustomAllreduce,
- )
+ custom_allreduce_kwargs = {
+ "group": self.cpu_group,
+ "device": self.device,
+ }
+ if current_platform.is_cuda():
+ from sglang.srt.distributed.device_communicators.custom_all_reduce import (
+ dispatch_custom_allreduce,
+ )
+ from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
+ CustomAllReduceV2,
+ )
- self.srt_custom_allreduce = CustomAllreduce(
- group=self.cpu_group,
- device=self.device,
- )
+ custom_allreduce_cls = dispatch_custom_allreduce(
+ group=self.cpu_group,
+ device=self.device,
+ )
+ if custom_allreduce_cls is CustomAllReduceV2:
+ custom_allreduce_kwargs["max_size"] = _DIFFUSION_CUSTOM_AR_MAX_SIZE
+ else:
+ # Preserve the existing ROCm and MUSA implementation selection.
+ from sglang.srt.distributed.device_communicators.custom_all_reduce import (
+ CustomAllreduce,
+ )
+
+ custom_allreduce_cls = CustomAllreduce
+
+ self.srt_custom_allreduce = custom_allreduce_cls(**custom_allreduce_kwargs)
@property
def first_rank(self):
@@ -374,9 +400,9 @@ class GroupCoordinator:
and not custom_ar.disabled
and custom_ar.should_custom_ar(input_)
):
- if custom_ar._IS_CAPTURING:
- return custom_ar.custom_all_reduce(input_)
- return custom_ar._all_reduce_impl(input_, registered=False)
+ output = custom_ar.custom_all_reduce(input_)
+ if output is not None:
+ return output
if (
current_platform.is_cpu()
and is_shm_available(input_.dtype, self.world_size, len(self.ranks))
diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py
index c25b8ff90..42506a6fd 100644
--- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py
+++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py
@@ -15,6 +15,7 @@ from sglang.kernels.ops.attention.flash_attention import flash_attn_varlen_func
from sglang.kernels.ops.diffusion import (
build_inv_indices,
fused_pack_qkv,
+ fused_pack_segmented_qkv,
fused_scatter_to_padded,
)
from sglang.multimodal_gen.runtime.breakable_cuda_graph.replay_token import (
@@ -827,6 +828,9 @@ class USPAttention(nn.Module):
attn_mask_meta: dict | None = None,
qkv_pre_all_to_all: bool = False,
seq_lens: list[int] | None = None,
+ q_prefix: torch.Tensor | None = None,
+ k_prefix: torch.Tensor | None = None,
+ v_prefix: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Forward pass for USPAttention.
@@ -889,6 +893,18 @@ class USPAttention(nn.Module):
if isinstance(attn_mask_meta, DynamicVarlenMaskMeta):
attn_mask_meta = attn_mask_meta.resolve(attn_mask)
+ segmented_prefix = q_prefix is not None
+ if segmented_prefix != (k_prefix is not None) or segmented_prefix != (
+ v_prefix is not None
+ ):
+ raise ValueError("q_prefix, k_prefix, and v_prefix must be set together")
+ if segmented_prefix and not (
+ effective_skip_sp or get_sequence_parallel_world_size() == 1
+ ):
+ raise NotImplementedError(
+ "Segmented QKV input currently supports only the local attention path."
+ )
+
# Tail-pad meta alone (sp_shard.tail_attn_meta; mask derivable from the
# pad span) also opts into the masked SP branch. gap_* = legacy alias.
meta_pad_start = meta_pad_end = None
@@ -991,9 +1007,20 @@ class USPAttention(nn.Module):
and q.device.type == "cuda"
and attn_mask.device == q.device
and q.dtype in (torch.float16, torch.bfloat16)
- and q.shape[:2] == attn_mask.shape == k.shape[:2] == v.shape[:2]
+ and (
+ (q.shape[0], q.shape[1] + q_prefix.shape[1])
+ if segmented_prefix
+ else q.shape[:2]
+ )
+ == attn_mask.shape
+ and q.shape == k.shape == v.shape
+ and (
+ not segmented_prefix
+ or q_prefix.shape == k_prefix.shape == v_prefix.shape
+ )
):
- bs, seq = q.shape[0], q.shape[1]
+ bs = q.shape[0]
+ seq = q.shape[1] + (q_prefix.shape[1] if segmented_prefix else 0)
indices = attn_mask_meta["indices"]
cu_seqlens = attn_mask_meta["cu_seqlens"]
max_seqlen = attn_mask_meta["max_seqlen"]
@@ -1008,7 +1035,46 @@ class USPAttention(nn.Module):
# (Joint attention with an image side is always non-empty
# in practice, so this only guards malformed inputs.)
if indices.shape[0] > 0:
- q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, k, v, indices)
+ all_valid = indices.shape[0] == bs * seq
+ if segmented_prefix:
+ q_unpad, k_unpad, v_unpad = fused_pack_segmented_qkv(
+ q_prefix,
+ k_prefix,
+ v_prefix,
+ q,
+ k,
+ v,
+ indices,
+ )
+ else:
+ if all_valid:
+ q_unpad, k_unpad, v_unpad = q, k, v
+ else:
+ q_unpad, k_unpad, v_unpad = fused_pack_qkv(
+ q, k, v, indices
+ )
+ if bs == 1 or all_valid:
+ # Empty cu_seqlens selects FA3's faster static
+ # persistent scheduler. A single packed sequence is
+ # dense even when its BCG bucket contains padding.
+ dense_seq = indices.shape[0] if bs == 1 else seq
+ out_dense = flash_attn_varlen_func(
+ q=q_unpad.reshape(bs, dense_seq, *q_unpad.shape[-2:]),
+ k=k_unpad.reshape(bs, dense_seq, *k_unpad.shape[-2:]),
+ v=v_unpad.reshape(bs, dense_seq, *v_unpad.shape[-2:]),
+ cu_seqlens_q=None,
+ cu_seqlens_k=None,
+ max_seqlen_q=dense_seq,
+ max_seqlen_k=dense_seq,
+ softmax_scale=self.softmax_scale,
+ causal=False,
+ ver=_fa_backend.fa_ver,
+ )
+ if all_valid:
+ return out_dense
+ return fused_scatter_to_padded(
+ out_dense.flatten(0, 1), inv_indices, bs, seq
+ )
out_unpad = flash_attn_varlen_func(
q=q_unpad,
k=k_unpad,
@@ -1023,6 +1089,11 @@ class USPAttention(nn.Module):
)
return fused_scatter_to_padded(out_unpad, inv_indices, bs, seq)
+ if segmented_prefix:
+ q = torch.cat([q_prefix, q], dim=1)
+ k = torch.cat([k_prefix, k], dim=1)
+ v = torch.cat([v_prefix, v], dim=1)
+
q_ = q.transpose(1, 2)
k_ = k.transpose(1, 2)
v_ = v.transpose(1, 2)
diff --git a/python/sglang/multimodal_gen/runtime/layers/linear.py b/python/sglang/multimodal_gen/runtime/layers/linear.py
index 2ed258277..5865b6896 100644
--- a/python/sglang/multimodal_gen/runtime/layers/linear.py
+++ b/python/sglang/multimodal_gen/runtime/layers/linear.py
@@ -126,6 +126,26 @@ class LinearMethodBase(QuantizeMethodBase):
raise NotImplementedError
+def apply_unquantized_linear(
+ x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor | None = None
+) -> torch.Tensor:
+ """Apply a plain linear projection with the runtime's reference semantics."""
+ if x.device.type == "mps":
+ if x.dtype == weight.dtype and (bias is None or bias.dtype == x.dtype):
+ return F.linear(x, weight, bias)
+ return F.linear(
+ x.to(torch.float32),
+ weight.to(torch.float32),
+ None if bias is None else bias.to(torch.float32),
+ ).to(x.dtype)
+
+ return (
+ F.linear(x, weight, bias)
+ if IS_AMP_SUPPORTED or bias is None
+ else F.linear(x, weight, bias.to(x.dtype))
+ )
+
+
class UnquantizedLinearMethod(LinearMethodBase):
"""Linear method without quantization."""
@@ -154,23 +174,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
def apply(
self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None
) -> torch.Tensor:
- if x.device.type == "mps":
- if x.dtype == layer.weight.dtype and (
- bias is None or bias.dtype == x.dtype
- ):
- return F.linear(x, layer.weight, bias)
- return F.linear(
- x.to(torch.float32),
- layer.weight.to(torch.float32),
- None if bias is None else bias.to(torch.float32),
- ).to(x.dtype)
-
- output = (
- F.linear(x, layer.weight, bias)
- if IS_AMP_SUPPORTED or bias is None
- else F.linear(x, layer.weight, bias.to(x.dtype))
- ) # NOTE: explicit dtype cast for bias is needed on platforms where amp isn't supported
- return output
+ return apply_unquantized_linear(x, layer.weight, bias)
class LinearBase(torch.nn.Module):
diff --git a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py
index e040f53e8..e0c14e77e 100644
--- a/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py
+++ b/python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py
@@ -1062,6 +1062,9 @@ def _resolve_quant_config(
if arch_config is None:
arch_config = server_args.pipeline_config.dit_config.arch_config
param_names_mapping_dict = arch_config.param_names_mapping
+ quant_param_names_mapping_dict = getattr(
+ arch_config, "quant_param_names_mapping", param_names_mapping_dict
+ )
reverse_param_names_mapping_dict = arch_config.reverse_param_names_mapping
quant_ignore_remap_dict = arch_config.quant_ignore_remap
@@ -1133,7 +1136,7 @@ def _resolve_quant_config(
fallback_group_size = getattr(quant_config, "group_size", None)
inferred_nvfp4_config = build_nvfp4_config_from_safetensors_list(
safetensors_list,
- param_names_mapping_dict,
+ quant_param_names_mapping_dict,
reverse_param_names_mapping_dict,
fallback_group_size,
)
diff --git a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py
index 3a6a97f05..be711a340 100644
--- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py
+++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py
@@ -24,6 +24,8 @@ from sglang.kernels.ops.diffusion import (
fused_linear_gelu_tanh,
is_plain_layer_norm,
mark_fused_gelu_site,
+ mark_qwen_image_added_qkv_site,
+ qwen_image_added_qkv_active,
try_fused_bias_mul_add,
try_fused_bias_scale_residual_norm_scale_shift,
try_fused_norm_scale_shift_fp8,
@@ -70,6 +72,7 @@ from sglang.multimodal_gen.runtime.layers.linear import (
MergedColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear,
+ apply_unquantized_linear,
)
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig,
@@ -194,6 +197,25 @@ def _local_seq_len(seq_len: int, sp_world_size: int) -> int:
_get_qkv_projections = get_qkv_projections
+def _split_unquantized_merged_linear(
+ linear: MergedColumnParallelLinear, x: torch.Tensor
+) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ """Apply a packed Q/K/V weight as three reference linear projections."""
+ sizes = linear.output_partition_sizes
+ if len(sizes) != 3:
+ raise ValueError(f"Expected three packed projection shards, got {sizes}")
+ weights = linear.weight.split(sizes, dim=0)
+ biases = (
+ linear.bias.split(sizes, dim=0)
+ if linear.bias is not None
+ else (None, None, None)
+ )
+ return tuple(
+ apply_unquantized_linear(x, weight, bias)
+ for weight, bias in zip(weights, biases)
+ )
+
+
def _can_defer_modelopt_output_bias(
quant_config: Optional[QuantizationConfig], capability: Any
) -> bool:
@@ -744,6 +766,7 @@ class QwenImageCrossAttention(nn.Module):
self.num_heads % tp_size == 0
), f"num_heads ({self.num_heads}) must be divisible by tp_size ({tp_size})"
self.local_num_heads = self.num_heads // tp_size
+ self._unquantized_added_qkv_is_packed = False
if self.use_fused_qkv:
# Use fused QKV projection for nunchaku quantization
@@ -785,8 +808,11 @@ class QwenImageCrossAttention(nn.Module):
self.norm_k = RMSNorm(head_dim, eps=eps) if qk_norm else nn.Identity()
if added_kv_proj_dim is not None:
+ self._unquantized_added_qkv_is_packed = quant_config is None
self.use_fused_added_qkv = (
- isinstance(quant_config, NunchakuConfig) or quant_name == "modelopt_fp8"
+ self._unquantized_added_qkv_is_packed
+ or isinstance(quant_config, NunchakuConfig)
+ or quant_name == "modelopt_fp8"
)
if self.use_fused_added_qkv:
self.to_added_qkv = MergedColumnParallelLinear(
@@ -796,6 +822,10 @@ class QwenImageCrossAttention(nn.Module):
quant_config=quant_config,
prefix=f"{prefix}.to_added_qkv",
)
+ if self._unquantized_added_qkv_is_packed:
+ # Packing changes BF16 GEMM reduction association. Keep it
+ # off for lossless requests and mount it for quality=high.
+ mark_qwen_image_added_qkv_site(self)
else:
self.add_q_proj = ColumnParallelLinear(
added_kv_proj_dim,
@@ -872,6 +902,25 @@ class QwenImageCrossAttention(nn.Module):
},
)
+ def _get_added_qkv_projections(
+ self, encoder_hidden_states: torch.Tensor
+ ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ if self.use_fused_added_qkv:
+ if (
+ self._unquantized_added_qkv_is_packed
+ and not qwen_image_added_qkv_active(self)
+ ):
+ return _split_unquantized_merged_linear(
+ self.to_added_qkv, encoder_hidden_states
+ )
+ added_qkv, _ = self.to_added_qkv(encoder_hidden_states)
+ return tuple(t.contiguous() for t in added_qkv.chunk(3, dim=-1))
+
+ encoder_query, _ = self.add_q_proj(encoder_hidden_states)
+ encoder_key, _ = self.add_k_proj(encoder_hidden_states)
+ encoder_value, _ = self.add_v_proj(encoder_hidden_states)
+ return encoder_query, encoder_key, encoder_value
+
def forward(
self,
hidden_states: torch.Tensor,
@@ -901,19 +950,31 @@ class QwenImageCrossAttention(nn.Module):
# Rows of tail padding inside THIS rank's text chunk (sp_shard meta).
sp_txt_pad = _attn_mask_meta_local_pad(attn_mask_meta)
- (
- img_query,
- img_key,
- img_value,
- txt_query,
- txt_key,
- txt_value,
- ) = _get_qkv_projections(
- self,
- hidden_states,
- encoder_hidden_states,
- make_contiguous=not self.use_fused_qkv_epilogue,
- )
+ if self._unquantized_added_qkv_is_packed and not qwen_image_added_qkv_active(
+ self
+ ):
+ img_query, img_key, img_value, _, _, _ = _get_qkv_projections(
+ self,
+ hidden_states,
+ make_contiguous=not self.use_fused_qkv_epilogue,
+ )
+ txt_query, txt_key, txt_value = self._get_added_qkv_projections(
+ encoder_hidden_states
+ )
+ else:
+ (
+ img_query,
+ img_key,
+ img_value,
+ txt_query,
+ txt_key,
+ txt_value,
+ ) = _get_qkv_projections(
+ self,
+ hidden_states,
+ encoder_hidden_states,
+ make_contiguous=not self.use_fused_qkv_epilogue,
+ )
# Reshape for multi-head attention
img_query = img_query.unflatten(-1, (self.local_num_heads, self.head_dim))
@@ -1001,6 +1062,17 @@ class QwenImageCrossAttention(nn.Module):
# Joint order [text, image]; join_seqs relocates any SP text tail-pad
# behind the image (see sp_shard.join_seqs for why).
+ if attn_mask is None and encoder_hidden_states_mask is not None:
+ image_mask = torch.ones(
+ (hidden_states.shape[0], img_query.shape[1]),
+ device=encoder_hidden_states_mask.device,
+ dtype=torch.bool,
+ )
+ attn_mask = torch.cat(
+ [encoder_hidden_states_mask.to(dtype=torch.bool), image_mask],
+ dim=1,
+ )
+
seg_qkv = None
# The segmented pre-all-to-all emits Ulysses layout; K/V-gather takes
# the join_seqs path and exchanges inside the attention instead.
@@ -1022,20 +1094,15 @@ class QwenImageCrossAttention(nn.Module):
joint_query, joint_key, joint_value = joint_qkv
elif seg_qkv is not None:
joint_query, joint_key, joint_value = seg_qkv
+ elif attn_mask is not None and not sp_text_sharded:
+ # Let the eager attention break point pack directly from the text
+ # and image segments. Materializing three dense joint tensors here
+ # only to gather their valid rows again wastes one launch per Q/K/V.
+ joint_query, joint_key, joint_value = img_query, img_key, img_value
else:
joint_query = join_seqs(txt_query, img_query, sp_txt_pad)
joint_key = join_seqs(txt_key, img_key, sp_txt_pad)
joint_value = join_seqs(txt_value, img_value, sp_txt_pad)
- if attn_mask is None and encoder_hidden_states_mask is not None:
- image_mask = torch.ones(
- (hidden_states.shape[0], img_query.shape[1]),
- device=encoder_hidden_states_mask.device,
- dtype=torch.bool,
- )
- attn_mask = torch.cat(
- [encoder_hidden_states_mask.to(dtype=torch.bool), image_mask],
- dim=1,
- )
# Compute joint attention
joint_hidden_states = self.attn(
@@ -1046,6 +1113,15 @@ class QwenImageCrossAttention(nn.Module):
attn_mask_meta=attn_mask_meta,
num_replicated_prefix=0 if sp_text_sharded else seq_len_txt,
qkv_pre_all_to_all=seg_qkv is not None,
+ q_prefix=(
+ txt_query if attn_mask is not None and not sp_text_sharded else None
+ ),
+ k_prefix=(
+ txt_key if attn_mask is not None and not sp_text_sharded else None
+ ),
+ v_prefix=(
+ txt_value if attn_mask is not None and not sp_text_sharded else None
+ ),
)
# Reshape back
diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
index 39b179915..fbda98243 100644
--- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
+++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py
@@ -28,6 +28,7 @@ from sglang.kernels.ops.diffusion import (
mount_lingbot_video_rmsnorm,
mount_ltx2_rms_norm_modulate,
mount_nvfp4_bias_gelu,
+ mount_qwen_image_added_qkv,
mount_sana_video_linear_attention,
unmount_fused_gate_rmsnorm,
unmount_fused_linear_gelu,
@@ -36,6 +37,7 @@ from sglang.kernels.ops.diffusion import (
unmount_lingbot_video_rmsnorm,
unmount_ltx2_rms_norm_modulate,
unmount_nvfp4_bias_gelu,
+ unmount_qwen_image_added_qkv,
unmount_sana_video_linear_attention,
)
from sglang.multimodal_gen import envs
@@ -173,6 +175,11 @@ _QUALITY_FUSION_HANDLERS: tuple[
mount_nvfp4_bias_gelu,
unmount_nvfp4_bias_gelu,
),
+ (
+ "Qwen-Image fused added-QKV",
+ mount_qwen_image_added_qkv,
+ unmount_qwen_image_added_qkv,
+ ),
(
"fused LN+modulate (affine folding)",
mount_fused_ln_modulate,
diff --git a/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_tp_graph_capture.py b/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_tp_graph_capture.py
index 7cf37e570..1ce940dea 100644
--- a/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_tp_graph_capture.py
+++ b/python/sglang/multimodal_gen/test/unit/test_diffusion_bcg_tp_graph_capture.py
@@ -31,6 +31,7 @@ from sglang.multimodal_gen.runtime.breakable_cuda_graph.runner import (
BaseBreakableCudaGraphRunner,
)
from sglang.multimodal_gen.runtime.distributed.group_coordinator import (
+ _DIFFUSION_CUSTOM_AR_MAX_SIZE,
GraphCaptureContext,
GroupCoordinator,
)
@@ -52,6 +53,67 @@ def _recording_context(events: list, name: str):
class TestBCGTPGraphCapture(CustomTestCase):
+ def test_cuda_custom_allreduce_uses_v2_dispatch_and_diffusion_workspace(self):
+ expected = object()
+ custom_allreduce_cls = MagicMock(return_value=expected)
+ group = SimpleNamespace(cpu_group=object(), device=torch.device("cuda:0"))
+
+ with patch(
+ "sglang.multimodal_gen.runtime.distributed.group_coordinator.current_platform.is_cuda",
+ return_value=True,
+ ), patch(
+ "sglang.srt.distributed.device_communicators.custom_all_reduce.dispatch_custom_allreduce",
+ return_value=custom_allreduce_cls,
+ ) as dispatch, patch(
+ "sglang.srt.distributed.device_communicators.custom_all_reduce_v2.CustomAllReduceV2",
+ custom_allreduce_cls,
+ ):
+ GroupCoordinator._init_srt_custom_allreduce(group)
+
+ dispatch.assert_called_once_with(group=group.cpu_group, device=group.device)
+ custom_allreduce_cls.assert_called_once_with(
+ group=group.cpu_group,
+ device=group.device,
+ max_size=_DIFFUSION_CUSTOM_AR_MAX_SIZE,
+ )
+ self.assertIs(group.srt_custom_allreduce, expected)
+
+ def test_non_cuda_custom_allreduce_preserves_default_workspace(self):
+ expected = object()
+ custom_allreduce_cls = MagicMock(return_value=expected)
+ group = SimpleNamespace(cpu_group=object(), device=object())
+
+ with patch(
+ "sglang.multimodal_gen.runtime.distributed.group_coordinator.current_platform.is_cuda",
+ return_value=False,
+ ), patch(
+ "sglang.srt.distributed.device_communicators.custom_all_reduce.CustomAllreduce",
+ custom_allreduce_cls,
+ ):
+ GroupCoordinator._init_srt_custom_allreduce(group)
+
+ custom_allreduce_cls.assert_called_once_with(
+ group=group.cpu_group,
+ device=group.device,
+ )
+ self.assertIs(group.srt_custom_allreduce, expected)
+
+ def test_all_reduce_uses_public_custom_allreduce_api(self):
+ output = object()
+ custom_ar = SimpleNamespace(
+ disabled=False,
+ should_custom_ar=MagicMock(return_value=True),
+ custom_all_reduce=MagicMock(return_value=output),
+ )
+ group = SimpleNamespace(world_size=2, srt_custom_allreduce=custom_ar)
+ input_ = SimpleNamespace(is_cpu=False)
+
+ result = GroupCoordinator.all_reduce(group, input_)
+
+ custom_ar.should_custom_ar.assert_called_once_with(input_)
+ custom_ar.custom_all_reduce.assert_called_once_with(input_)
+ self.assertIs(result, output)
+
# --- GroupCoordinator.graph_capture -> CustomAllreduce.capture ---------- #
def _run_graph_capture(self, custom_ar, events):
diff --git a/test/registered/kernels/benchmark/diffusion/bench_varlen_segmented_pack.py b/test/registered/kernels/benchmark/diffusion/bench_varlen_segmented_pack.py
new file mode 100644
index 000000000..3c55ba019
--- /dev/null
+++ b/test/registered/kernels/benchmark/diffusion/bench_varlen_segmented_pack.py
@@ -0,0 +1,95 @@
+import torch
+
+from sglang.kernels.jit.benchmark import marker
+from sglang.kernels.ops.diffusion import fused_pack_qkv, fused_pack_segmented_qkv
+from sglang.test.ci.ci_register import register_cuda_ci
+
+register_cuda_ci(
+ est_time=12, stage="base-b-kernel-benchmark", runner_config="1-gpu-large"
+)
+
+DEVICE = "cuda"
+DTYPE = torch.bfloat16
+
+
+def _materialized_pack(q_prefix, k_prefix, v_prefix, q_main, k_main, v_main, indices):
+ return fused_pack_qkv(
+ torch.cat([q_prefix, q_main], dim=1),
+ torch.cat([k_prefix, k_main], dim=1),
+ torch.cat([v_prefix, v_main], dim=1),
+ indices,
+ )
+
+
+def _segmented_pack(q_prefix, k_prefix, v_prefix, q_main, k_main, v_main, indices):
+ return fused_pack_segmented_qkv(
+ q_prefix, k_prefix, v_prefix, q_main, k_main, v_main, indices
+ )
+
+
+@marker.parametrize(
+ "batch,prefix_rows,main_rows,heads,head_dim,valid_prefix_rows",
+ [
+ (1, 64, 4096, 12, 128, 24),
+ (2, 64, 1024, 8, 128, 40),
+ ],
+ ci_vals=[(1, 64, 4096, 12, 128, 24)],
+)
+@marker.benchmark("provider", ["materialized", "segmented"])
+def benchmark(
+ batch: int,
+ prefix_rows: int,
+ main_rows: int,
+ heads: int,
+ head_dim: int,
+ valid_prefix_rows: int,
+ provider: str,
+) -> marker.BenchResult:
+ generator = torch.Generator(device=DEVICE).manual_seed(42)
+ prefixes = tuple(
+ torch.randn(
+ batch,
+ prefix_rows,
+ heads,
+ head_dim,
+ dtype=DTYPE,
+ device=DEVICE,
+ generator=generator,
+ )
+ for _ in range(3)
+ )
+ mains = tuple(
+ torch.randn(
+ batch,
+ main_rows,
+ heads,
+ head_dim,
+ dtype=DTYPE,
+ device=DEVICE,
+ generator=generator,
+ )
+ for _ in range(3)
+ )
+ mask = torch.zeros(batch, prefix_rows + main_rows, dtype=torch.bool, device=DEVICE)
+ mask[:, :valid_prefix_rows] = True
+ mask[:, prefix_rows:] = True
+ indices = mask.flatten().nonzero(as_tuple=False).flatten()
+ args = (*prefixes, *mains, indices)
+
+ expected = _materialized_pack(*args)
+ actual = _segmented_pack(*args)
+ assert all(
+ torch.equal(got, want) for got, want in zip(actual, expected, strict=True)
+ )
+
+ fn = _materialized_pack if provider == "materialized" else _segmented_pack
+ return marker.do_bench(
+ fn,
+ input_args=args,
+ graph_clone_args=tuple(range(len(args))),
+ disable_log_bandwidth=True,
+ )
+
+
+if __name__ == "__main__":
+ benchmark.run()
diff --git a/test/registered/kernels/ops/diffusion/test_layout.py b/test/registered/kernels/ops/diffusion/test_layout.py
index 9cf834bba..ff8bdafb6 100644
--- a/test/registered/kernels/ops/diffusion/test_layout.py
+++ b/test/registered/kernels/ops/diffusion/test_layout.py
@@ -31,6 +31,7 @@ from sglang.kernels.ops.diffusion import (
from sglang.kernels.ops.diffusion import (
fused_causal_conv3d_cat_pad_cuda,
fused_pack_qkv,
+ fused_pack_segmented_qkv,
fused_scatter_to_padded,
pack_qkv_destination_major,
usp_merge_heads,
@@ -211,6 +212,28 @@ def test_varlen_pack_matches_index_select(dtype, shape):
assert torch.equal(got, want)
+@pytest.mark.parametrize("dtype", VARLEN_DTYPES)
+@pytest.mark.parametrize("shape", VARLEN_SHAPES, ids=lambda s: s[0])
+def test_varlen_segmented_pack_matches_materialized_joint(dtype, shape):
+ _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape
+ torch.manual_seed(42)
+ indices, _ = _build_meta(_build_mask(bs, s_txt, s_img, valid_txt_lens))
+ txt_qkv = tuple(
+ torch.randn(bs, s_txt, num_heads, head_dim, dtype=dtype, device=DEVICE)
+ for _ in range(3)
+ )
+ img_qkv = tuple(
+ torch.randn(bs, s_img, num_heads, head_dim, dtype=dtype, device=DEVICE)
+ for _ in range(3)
+ )
+
+ got = fused_pack_segmented_qkv(*txt_qkv, *img_qkv, indices)
+ for actual, txt, img in zip(got, txt_qkv, img_qkv, strict=True):
+ joint = torch.cat([txt, img], dim=1)
+ expected = joint.flatten(0, 1).index_select(0, indices)
+ assert torch.equal(actual, expected)
+
+
@pytest.mark.parametrize("dtype", VARLEN_DTYPES)
@pytest.mark.parametrize("shape", VARLEN_SHAPES, ids=lambda s: s[0])
def test_varlen_scatter_matches_index_copy(dtype, shape):
@@ -344,6 +367,45 @@ def _varlen_path(q, k, v, key_mask, softmax_scale):
return fused_scatter_to_padded(out_unpad, meta["inv_indices"], bs, seq)
+@pytest.mark.parametrize("dtype", VARLEN_DTYPES)
+def test_fa_dense_scheduler_matches_single_sequence_varlen(dtype):
+ torch.manual_seed(7)
+ batch_size, seq, num_heads, head_dim = 1, 256, 4, 128
+ q, k, v = (
+ torch.randn(
+ batch_size,
+ seq,
+ num_heads,
+ head_dim,
+ dtype=dtype,
+ device=DEVICE,
+ )
+ for _ in range(3)
+ )
+ cu_seqlens = torch.tensor([0, seq], dtype=torch.int32, device=DEVICE)
+ kwargs = dict(
+ max_seqlen_q=seq,
+ max_seqlen_k=seq,
+ softmax_scale=head_dim**-0.5,
+ causal=False,
+ ver=_fa_backend.fa_ver,
+ )
+ try:
+ varlen = flash_attn_varlen_func(
+ q.flatten(0, 1),
+ k.flatten(0, 1),
+ v.flatten(0, 1),
+ cu_seqlens,
+ cu_seqlens,
+ **kwargs,
+ ).view_as(q)
+ dense = flash_attn_varlen_func(q, k, v, None, None, **kwargs)
+ except ImportError as exc: # pragma: no cover - image-dependent
+ pytest.skip(f"FlashAttention unavailable: {exc}")
+
+ torch.testing.assert_close(dense, varlen, rtol=1e-3, atol=1e-3)
+
+
@pytest.mark.parametrize("dtype", VARLEN_DTYPES)
@pytest.mark.parametrize("shape", VARLEN_SHAPES, ids=lambda s: s[0])
def test_varlen_path_matches_sdpa_on_valid_rows(dtype, shape):
diff --git a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py
index ab7b2509e..0335b15c3 100644
--- a/test/registered/kernels/ops/diffusion/test_model_fast_paths.py
+++ b/test/registered/kernels/ops/diffusion/test_model_fast_paths.py
@@ -48,12 +48,15 @@ from sglang.kernels.ops.diffusion import (
mark_fused_ln_modulate_site,
mark_hunyuan_qknorm_site,
mark_ltx2_rms_norm_modulate_site,
+ mark_qwen_image_added_qkv_site,
mount_fused_ln_modulate,
mount_hunyuan_qknorm,
mount_ltx2_rms_norm_modulate,
+ mount_qwen_image_added_qkv,
try_flux2_token_cat_nvfp4,
unmount_hunyuan_qknorm,
unmount_ltx2_rms_norm_modulate,
+ unmount_qwen_image_added_qkv,
wan_rmsnorm_silu,
)
from sglang.kernels.ops.diffusion.common.platform import is_cuda
@@ -98,9 +101,11 @@ from sglang.multimodal_gen.runtime.models.dits.longcat_image import (
)
from sglang.multimodal_gen.runtime.models.dits.ltx_2 import _ltx2_rms_norm_modulate
from sglang.multimodal_gen.runtime.models.dits.qwen_image import (
+ QwenImageCrossAttention,
QwenImageTransformerBlock,
_qwen_modulation_cache_key,
_qwen_norm_out,
+ _split_unquantized_merged_linear,
)
from sglang.multimodal_gen.runtime.models.dits.sana import (
_eager_ln_modulate as _sana_eager_ln_modulate,
@@ -412,6 +417,64 @@ class TestFlux2EagerFusions(CustomTestCase):
# -------------------------------------------------------------------------
+class _PackedAddedQKV(nn.Module):
+ def __init__(self, dim: int):
+ super().__init__()
+ self.output_partition_sizes = [dim, dim, dim]
+ self.quant_config = None
+ self.weight = nn.Parameter(
+ torch.randn(3 * dim, dim, device="cuda", dtype=torch.bfloat16),
+ requires_grad=False,
+ )
+ self.bias = nn.Parameter(
+ torch.randn(3 * dim, device="cuda", dtype=torch.bfloat16),
+ requires_grad=False,
+ )
+ self.calls = 0
+
+ def forward(self, x):
+ self.calls += 1
+ return F.linear(x, self.weight, self.bias), None
+
+
+def test_qwen_added_qkv_lossless_uses_three_reference_gemms():
+ torch.manual_seed(20260831)
+ dim = 64
+ x = torch.randn(1, 17, dim, device="cuda", dtype=torch.bfloat16)
+ packed = _PackedAddedQKV(dim)
+
+ attention = QwenImageCrossAttention.__new__(QwenImageCrossAttention)
+ nn.Module.__init__(attention)
+ attention.use_fused_added_qkv = True
+ attention._unquantized_added_qkv_is_packed = True
+ attention.to_added_qkv = packed
+ mark_qwen_image_added_qkv_site(attention)
+
+ expected_lossless = _split_unquantized_merged_linear(packed, x)
+ actual_lossless = attention._get_added_qkv_projections(x)
+ assert packed.calls == 0
+ assert all(
+ torch.equal(actual, expected)
+ for actual, expected in zip(actual_lossless, expected_lossless)
+ )
+
+ assert mount_qwen_image_added_qkv(attention)
+ actual_high = attention._get_added_qkv_projections(x)
+ expected_high = tuple(
+ tensor.contiguous()
+ for tensor in F.linear(x, packed.weight, packed.bias).chunk(3, dim=-1)
+ )
+ assert packed.calls == 1
+ assert all(
+ torch.equal(actual, expected)
+ for actual, expected in zip(actual_high, expected_high)
+ )
+
+ unmount_qwen_image_added_qkv(attention)
+ attention._get_added_qkv_projections(x)
+ assert packed.calls == 1
+
+
class _CountingProjection(nn.Module):
def __init__(self, offset: float):
super().__init__()
diff --git a/test/registered/kernels/ops/diffusion/test_sites.py b/test/registered/kernels/ops/diffusion/test_sites.py
index ed71093e1..b00e90d3d 100644
--- a/test/registered/kernels/ops/diffusion/test_sites.py
+++ b/test/registered/kernels/ops/diffusion/test_sites.py
@@ -28,6 +28,7 @@ import torch.nn.functional as F
import sglang.kernels.ops.diffusion.sites.fused_gate_rmsnorm_site as gate_rmsnorm
import sglang.kernels.ops.diffusion.sites.fused_linear_gelu_site as linear_gelu
import sglang.kernels.ops.diffusion.sites.lingbot_video_rmsnorm_site as lingbot_video_rmsnorm
+import sglang.kernels.ops.diffusion.sites.qwen_image_added_qkv_site as qwen_image_added_qkv
import sglang.kernels.ops.diffusion.sites.sana_video_linear_attention_site as sana_video_linear_attention
from sglang.kernels.ops.diffusion import (
BitExactFusionGate,
@@ -91,6 +92,20 @@ def test_quality_gate_rejection_is_all_or_nothing():
assert not fusion.mount(nn.Module())
+def test_qwen_image_added_qkv_site_is_request_scoped():
+ site = nn.Module()
+ site.to_added_qkv = nn.Module()
+ site.to_added_qkv.quant_config = None
+ site.to_added_qkv.output_partition_sizes = [8, 8, 8]
+ qwen_image_added_qkv.mark_qwen_image_added_qkv_site(site)
+
+ assert not qwen_image_added_qkv.qwen_image_added_qkv_active(site)
+ assert qwen_image_added_qkv.mount_qwen_image_added_qkv(site)
+ assert qwen_image_added_qkv.qwen_image_added_qkv_active(site)
+ qwen_image_added_qkv.unmount_qwen_image_added_qkv(site)
+ assert not qwen_image_added_qkv.qwen_image_added_qkv_active(site)
+
+
# -------------------------------------------------------------------------
# BitExactFusionGate protocol (CPU)
# -------------------------------------------------------------------------