[Diffusion] Optimize Qwen-Image TP collectives and attention (#36680)
Co-authored-by: Mick <mickjagger19@icloud.com> Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
co-authored by
Mick
Cursor
parent
562b661e0e
commit
71cee04ebe
@@ -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).
|
||||
|
||||
|
||||
@@ -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>
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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`.
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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)
|
||||
@@ -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. |
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
@@ -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):
|
||||
|
||||
@@ -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__()
|
||||
|
||||
@@ -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)
|
||||
# -------------------------------------------------------------------------
|
||||
|
||||
Reference in New Issue
Block a user