From d616b8edad9aa2f8adfe7438b600cf46e9739a01 Mon Sep 17 00:00:00 2001 From: mispa-ms <81828223+mispa-ms@users.noreply.github.com> Date: Thu, 28 May 2026 06:26:09 -0700 Subject: [PATCH] [diffusion][jit_kernel] perf: varlen FA fast path for USPAttention masked branch (#26318) --- .../diffusion/triton/varlen_pack_pad.py | 191 +++++++++++++++++ .../tests/diffusion/test_varlen_pack_pad.py | 193 ++++++++++++++++++ .../test_varlen_uspattn_equivalence.py | 156 ++++++++++++++ .../runtime/layers/attention/__init__.py | 2 + .../runtime/layers/attention/layer.py | 93 +++++++++ .../runtime/models/dits/qwen_image.py | 17 +- 6 files changed, 649 insertions(+), 3 deletions(-) create mode 100644 python/sglang/jit_kernel/diffusion/triton/varlen_pack_pad.py create mode 100644 python/sglang/jit_kernel/tests/diffusion/test_varlen_pack_pad.py create mode 100644 python/sglang/jit_kernel/tests/diffusion/test_varlen_uspattn_equivalence.py diff --git a/python/sglang/jit_kernel/diffusion/triton/varlen_pack_pad.py b/python/sglang/jit_kernel/diffusion/triton/varlen_pack_pad.py new file mode 100644 index 000000000..937c3ad03 --- /dev/null +++ b/python/sglang/jit_kernel/diffusion/triton/varlen_pack_pad.py @@ -0,0 +1,191 @@ +"""Fused Triton pack/scatter kernels for the varlen mask path. + +Used by ``USPAttention.forward`` masked branch to gather Q/K/V at valid +positions and scatter the FA output back to the dense ``[B, S, H, D]`` layout. +""" + +from __future__ import annotations + +import torch +import triton # type: ignore +import triton.language as tl # type: ignore + +# --------------------------------------------------------------------------- +# Pack (unpad) — gather Q/K/V at indices into packed [total_valid, H, D] +# --------------------------------------------------------------------------- + + +@triton.jit +def _fused_pack_qkv_kernel( + Q_ptr, + K_ptr, + V_ptr, + Q_unpad_ptr, + K_unpad_ptr, + V_unpad_ptr, + indices_ptr, + HD, # H * D, flattened feature dim + src_row_stride, # stride between rows in Q/K/V (B*S row -> next row) + dst_row_stride, # stride in Q_unpad/K_unpad/V_unpad + BLOCK_HD: tl.constexpr, +): + """One program per packed row; copies Q[src], K[src], V[src] to dst row.""" + out_row = tl.program_id(0) + src_row = tl.load(indices_ptr + out_row).to(tl.int64) + + cols = tl.arange(0, BLOCK_HD) + col_mask = cols < HD + + src_offset = src_row * src_row_stride + cols + dst_offset = out_row * dst_row_stride + cols + + q_val = tl.load(Q_ptr + src_offset, mask=col_mask) + k_val = tl.load(K_ptr + src_offset, mask=col_mask) + v_val = tl.load(V_ptr + src_offset, mask=col_mask) + + 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_qkv( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + indices: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Pack ``[B, S, H, D]`` Q/K/V at ``indices`` into ``[total_valid, H, D]``. + + ``indices`` is the int64 flat ``B*S`` position for each kept token. + Non-contiguous inputs are made contiguous internally. + """ + assert q.shape == k.shape == v.shape, "Q/K/V must share shape" + assert q.dtype == k.dtype == v.dtype, "Q/K/V must share dtype" + assert q.dim() == 4, "Q/K/V must be [B, S, H, D]" + assert indices.dtype in (torch.int32, torch.int64) + q = q.contiguous() + k = k.contiguous() + v = v.contiguous() + bs, seq, num_heads, head_dim = q.shape + hd = num_heads * head_dim + n_valid = indices.shape[0] + + if n_valid == 0: + return ( + q.new_empty(0, num_heads, head_dim), + k.new_empty(0, num_heads, head_dim), + v.new_empty(0, num_heads, head_dim), + ) + + block_hd = triton.next_power_of_2(hd) + q_flat = q.view(bs * seq, hd) + k_flat = k.view(bs * seq, hd) + v_flat = v.view(bs * seq, hd) + q_unpad = torch.empty(n_valid, hd, dtype=q.dtype, device=q.device) + k_unpad = torch.empty(n_valid, hd, dtype=k.dtype, device=k.device) + v_unpad = torch.empty(n_valid, hd, dtype=v.dtype, device=v.device) + + with torch.get_device_module().device(q.device): + _fused_pack_qkv_kernel[(n_valid,)]( + q_flat, + k_flat, + v_flat, + q_unpad, + k_unpad, + v_unpad, + indices, + hd, + q_flat.stride(0), + q_unpad.stride(0), + BLOCK_HD=block_hd, + ) + + return ( + q_unpad.view(n_valid, num_heads, head_dim), + k_unpad.view(n_valid, num_heads, head_dim), + v_unpad.view(n_valid, num_heads, head_dim), + ) + + +# --------------------------------------------------------------------------- +# Scatter (pad) — write packed output to [B, S, H, D] with zeros at invalid +# --------------------------------------------------------------------------- + + +@triton.jit +def _fused_scatter_to_padded_kernel( + Out_unpad_ptr, + Out_padded_ptr, + inv_indices_ptr, # [B*S]: pack idx for valid row, -1 for invalid + HD, + src_row_stride, + dst_row_stride, + BLOCK_HD: tl.constexpr, +): + """One program per padded row; writes from pack or zeros.""" + pad_row = tl.program_id(0) + inv_idx = tl.load(inv_indices_ptr + pad_row).to(tl.int64) + + cols = tl.arange(0, BLOCK_HD) + col_mask = cols < HD + valid = inv_idx >= 0 + + safe_idx = tl.where(valid, inv_idx, 0) + src_offset = safe_idx * src_row_stride + cols + dst_offset = pad_row * dst_row_stride + cols + + val = tl.load(Out_unpad_ptr + src_offset, mask=col_mask & valid, other=0.0) + tl.store(Out_padded_ptr + dst_offset, val, mask=col_mask) + + +def fused_scatter_to_padded( + out_unpad: torch.Tensor, + inv_indices: torch.Tensor, + batch_size: int, + seqlen: int, +) -> torch.Tensor: + """Scatter packed varlen output back to ``[B, S, H, D]`` with zero padding. + + ``inv_indices`` is ``[B*S]`` giving the pack row index for each padded + position (``-1`` for padding). Non-contiguous ``out_unpad`` is made contiguous. + """ + assert out_unpad.dim() == 3, "out_unpad must be [total_valid, H, D]" + assert inv_indices.shape == (batch_size * seqlen,) + assert inv_indices.dtype in (torch.int32, torch.int64) + out_unpad = out_unpad.contiguous() + _, num_heads, head_dim = out_unpad.shape + hd = num_heads * head_dim + block_hd = triton.next_power_of_2(hd) + + out_padded = torch.empty( + batch_size * seqlen, hd, dtype=out_unpad.dtype, device=out_unpad.device + ) + out_unpad_flat = out_unpad.view(-1, hd) + + with torch.get_device_module().device(out_unpad.device): + _fused_scatter_to_padded_kernel[(batch_size * seqlen,)]( + out_unpad_flat, + out_padded, + inv_indices, + hd, + out_unpad_flat.stride(0), + out_padded.stride(0), + BLOCK_HD=block_hd, + ) + + return out_padded.view(batch_size, seqlen, num_heads, head_dim) + + +# --------------------------------------------------------------------------- +# Inverse-index builder (called once per request alongside indices) +# --------------------------------------------------------------------------- + + +def build_inv_indices(indices: torch.Tensor, total_rows: int) -> torch.Tensor: + """For each padded row in ``[B*S]``, return its pack index or ``-1``.""" + n_valid = indices.shape[0] + inv = torch.full((total_rows,), -1, dtype=torch.int32, device=indices.device) + inv[indices.long()] = torch.arange( + n_valid, dtype=torch.int32, device=indices.device + ) + return inv diff --git a/python/sglang/jit_kernel/tests/diffusion/test_varlen_pack_pad.py b/python/sglang/jit_kernel/tests/diffusion/test_varlen_pack_pad.py new file mode 100644 index 000000000..d25ba55d9 --- /dev/null +++ b/python/sglang/jit_kernel/tests/diffusion/test_varlen_pack_pad.py @@ -0,0 +1,193 @@ +"""Numerical correctness for fused varlen pack/scatter Triton kernels. + +Bit-exact comparison against the equivalent PyTorch ops (index_select, +zeros + index_copy_) across bf16/fp16 and several shape/mask cases. +""" + +import pytest +import torch + +from sglang.jit_kernel.diffusion.triton.varlen_pack_pad import ( + build_inv_indices, + fused_pack_qkv, + fused_scatter_to_padded, +) +from sglang.jit_kernel.utils import get_ci_test_range +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=10, suite="base-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=60, suite="nightly-kernel-1-gpu", nightly=True) + +DEVICE = "cuda" +DTYPES = get_ci_test_range([torch.bfloat16, torch.float16], [torch.bfloat16]) +# (bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens) tuples +SHAPES = get_ci_test_range( + [ + # name, bs, s_txt, s_img, H, D, valid_txt_lens + ("small_c2", 2, 64, 128, 4, 64, [32, 48]), + ("prod_c2", 2, 256, 1024, 24, 128, [128, 200]), + ("all_valid_b1", 1, 64, 128, 4, 64, [64]), + ("all_valid_b4", 4, 64, 128, 4, 64, [64, 64, 64, 64]), + ("c8_prod", 8, 256, 4096, 24, 128, [128, 200, 256, 100, 50, 256, 256, 50]), + # one batch with zero valid text tokens (image side still valid) + ("zero_txt_one_batch", 2, 64, 128, 4, 64, [0, 32]), + # bs=1 with no text validity (only image rows packed) + ("bs1_zero_txt", 1, 64, 128, 4, 64, [0]), + ], + [ + ("small_c2", 2, 64, 128, 4, 64, [32, 48]), + ("prod_c2", 2, 256, 1024, 24, 128, [128, 200]), + ("all_valid_b4", 4, 64, 128, 4, 64, [64, 64, 64, 64]), + ], +) + + +def _build_mask(bs, s_txt, s_img, valid_txt_lens): + s = s_txt + s_img + mask = torch.zeros(bs, s, dtype=torch.bool, device=DEVICE) + for b, vt in enumerate(valid_txt_lens): + mask[b, :vt] = True + mask[b, s_txt:] = True + return mask + + +def _ref_pack(q, k, v, indices): + bs, seq = q.shape[:2] + flat = lambda t: t.reshape(bs * seq, *t.shape[2:]) + return ( + flat(q).index_select(0, indices), + flat(k).index_select(0, indices), + flat(v).index_select(0, indices), + ) + + +def _ref_scatter(out_unpad, indices, bs, seq): + n_valid = indices.shape[0] + _, num_heads, head_dim = out_unpad.shape + flat = torch.zeros( + bs * seq, num_heads, head_dim, dtype=out_unpad.dtype, device=DEVICE + ) + flat.index_copy_(0, indices, out_unpad) + return flat.view(bs, seq, num_heads, head_dim) + + +def _build_meta(mask): + bs, seq = mask.shape + indices = mask.reshape(-1).nonzero(as_tuple=False).flatten() + inv_indices = build_inv_indices(indices, bs * seq) + return indices, inv_indices + + +@pytest.mark.parametrize("dtype", DTYPES) +@pytest.mark.parametrize( + "shape", SHAPES, ids=lambda s: s[0] if isinstance(s, tuple) else str(s) +) +def test_pack_matches_index_select(dtype, shape): + _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape + torch.manual_seed(0) + s = s_txt + s_img + mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) + indices, _ = _build_meta(mask) + + q = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + k = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + v = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + + q_ref, k_ref, v_ref = _ref_pack(q, k, v, indices) + q_fused, k_fused, v_fused = fused_pack_qkv(q, k, v, indices) + + # bit-exact: pack is pure gather, no math + assert torch.equal(q_ref, q_fused) + assert torch.equal(k_ref, k_fused) + assert torch.equal(v_ref, v_fused) + + +@pytest.mark.parametrize("dtype", DTYPES) +@pytest.mark.parametrize( + "shape", SHAPES, ids=lambda s: s[0] if isinstance(s, tuple) else str(s) +) +def test_scatter_matches_index_copy(dtype, shape): + _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape + torch.manual_seed(1) + s = s_txt + s_img + mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) + indices, inv_indices = _build_meta(mask) + n_valid = indices.shape[0] + + out_unpad = torch.randn(n_valid, num_heads, head_dim, dtype=dtype, device=DEVICE) + out_ref = _ref_scatter(out_unpad, indices, bs, s) + out_fused = fused_scatter_to_padded(out_unpad, inv_indices, bs, s) + + # bit-exact: scatter is pure copy + zero-fill + assert torch.equal(out_ref, out_fused) + # Padding rows must be exactly zero + invalid = ~mask + if invalid.any(): + assert out_fused[invalid].abs().max().item() == 0.0 + + +def test_pack_handles_non_contiguous_input(): + """Helper must accept non-contiguous Q/K/V (auto .contiguous() inside).""" + torch.manual_seed(2) + bs, s_txt, s_img, num_heads, head_dim = 2, 64, 128, 4, 64 + s = s_txt + s_img + mask = _build_mask(bs, s_txt, s_img, [32, 48]) + indices, _ = _build_meta(mask) + + # Build non-contiguous tensors via permute + qkv_pre = torch.randn( + bs, num_heads, s, head_dim, dtype=torch.bfloat16, device=DEVICE + ) + q = qkv_pre.permute(0, 2, 1, 3) + k = torch.randn_like(qkv_pre).permute(0, 2, 1, 3) + v = torch.randn_like(qkv_pre).permute(0, 2, 1, 3) + assert not q.is_contiguous() + + q_ref, k_ref, v_ref = _ref_pack( + q.contiguous(), k.contiguous(), v.contiguous(), indices + ) + q_fused, k_fused, v_fused = fused_pack_qkv(q, k, v, indices) + assert torch.equal(q_ref, q_fused) + assert torch.equal(k_ref, k_fused) + assert torch.equal(v_ref, v_fused) + + +def test_build_inv_indices_matches_manual(): + """build_inv_indices output should match the manual full+scatter form.""" + torch.manual_seed(3) + bs, s = 2, 32 + mask = torch.bernoulli(torch.full((bs, s), 0.6, device=DEVICE)).to(torch.bool) + indices = mask.reshape(-1).nonzero(as_tuple=False).flatten() + n_valid = indices.shape[0] + + manual = torch.full((bs * s,), -1, dtype=torch.int32, device=DEVICE) + if n_valid > 0: + manual[indices.long()] = torch.arange(n_valid, dtype=torch.int32, device=DEVICE) + + built = build_inv_indices(indices, bs * s) + assert torch.equal(built, manual) + + +def test_empty_valid_set_handled(): + """All-False mask: pack returns empty tensors; scatter writes all zeros.""" + bs, s, num_heads, head_dim = 2, 16, 4, 64 + mask = torch.zeros(bs, s, dtype=torch.bool, device=DEVICE) + indices = mask.reshape(-1).nonzero(as_tuple=False).flatten() + inv_indices = build_inv_indices(indices, bs * s) + assert indices.numel() == 0 + + q = torch.randn(bs, s, num_heads, head_dim, dtype=torch.bfloat16, device=DEVICE) + q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, q.clone(), q.clone(), indices) + assert q_unpad.shape == (0, num_heads, head_dim) + assert k_unpad.shape == (0, num_heads, head_dim) + assert v_unpad.shape == (0, num_heads, head_dim) + + out_padded = fused_scatter_to_padded(q_unpad, inv_indices, bs, s) + assert out_padded.shape == (bs, s, num_heads, head_dim) + assert out_padded.abs().max().item() == 0.0 + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/python/sglang/jit_kernel/tests/diffusion/test_varlen_uspattn_equivalence.py b/python/sglang/jit_kernel/tests/diffusion/test_varlen_uspattn_equivalence.py new file mode 100644 index 000000000..ef45abef5 --- /dev/null +++ b/python/sglang/jit_kernel/tests/diffusion/test_varlen_uspattn_equivalence.py @@ -0,0 +1,156 @@ +"""End-to-end equivalence between USPAttention varlen path and SDPA reference. + +Compares the production varlen path (``build_varlen_mask_meta`` + +``fused_pack_qkv`` + ``flash_attn_varlen_func`` + ``fused_scatter_to_padded``) +against ``torch.nn.functional.scaled_dot_product_attention`` with a broadcast +key mask, for inputs the gating in ``USPAttention.forward`` would accept. + +Verifies the documented contract: + * Valid (non-masked) query rows match SDPA within FA-vs-SDPA tolerance. + * Masked query rows are exactly zero in the varlen path (differs from + SDPA, which produces deterministic attention output at those rows). +""" + +import pytest +import torch +import torch.nn.functional as F + +from sglang.jit_kernel.diffusion.triton.varlen_pack_pad import ( + fused_pack_qkv, + fused_scatter_to_padded, +) +from sglang.jit_kernel.flash_attention import flash_attn_varlen_func +from sglang.jit_kernel.utils import get_ci_test_range +from sglang.multimodal_gen.runtime.layers.attention.backends import ( + flash_attn as _fa_backend, +) +from sglang.multimodal_gen.runtime.layers.attention.layer import ( + build_varlen_mask_meta, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=15, suite="base-b-kernel-unit-1-gpu-large") +register_cuda_ci(est_time=60, suite="nightly-kernel-1-gpu", nightly=True) + +DEVICE = "cuda" +DTYPES = get_ci_test_range([torch.bfloat16, torch.float16], [torch.bfloat16]) +# (name, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens) +SHAPES = get_ci_test_range( + [ + ("small_c2", 2, 64, 128, 4, 64, [32, 48]), + ("prod_c2", 2, 256, 1024, 24, 128, [128, 200]), + ("all_valid_b1", 1, 64, 128, 4, 64, [64]), + ("zero_txt_one_batch", 2, 64, 128, 4, 64, [0, 32]), + ], + [ + ("small_c2", 2, 64, 128, 4, 64, [32, 48]), + ], +) + + +def _build_mask(bs, s_txt, s_img, valid_txt_lens): + s = s_txt + s_img + mask = torch.zeros(bs, s, dtype=torch.bool, device=DEVICE) + for b, vt in enumerate(valid_txt_lens): + mask[b, :vt] = True + mask[b, s_txt:] = True + return mask + + +def _sdpa_with_key_mask(q, k, v, key_mask, softmax_scale): + """Reference: SDPA with a ``[B, S]`` key mask broadcast to ``[B, 1, 1, S]``.""" + q_ = q.transpose(1, 2) + k_ = k.transpose(1, 2) + v_ = v.transpose(1, 2) + mask = key_mask.to(dtype=q.dtype)[:, None, None, :] + mask = (mask - 1.0) * torch.finfo(q.dtype).max + out = F.scaled_dot_product_attention( + q_, + k_, + v_, + attn_mask=mask, + dropout_p=0.0, + is_causal=False, + scale=softmax_scale, + ) + return out.transpose(1, 2) + + +def _varlen_path(q, k, v, key_mask, softmax_scale): + """Production varlen path matching USPAttention.forward masked branch.""" + bs, seq = q.shape[0], q.shape[1] + meta = build_varlen_mask_meta(key_mask) + indices = meta["indices"] + if indices.shape[0] == 0: + return torch.zeros_like(q) + q_unpad, k_unpad, v_unpad = fused_pack_qkv(q, k, v, indices) + out_unpad = flash_attn_varlen_func( + q=q_unpad, + k=k_unpad, + v=v_unpad, + cu_seqlens_q=meta["cu_seqlens"], + cu_seqlens_k=meta["cu_seqlens"], + max_seqlen_q=meta["max_seqlen"], + max_seqlen_k=meta["max_seqlen"], + softmax_scale=softmax_scale, + causal=False, + ver=_fa_backend.fa_ver, + ) + return fused_scatter_to_padded(out_unpad, meta["inv_indices"], bs, seq) + + +@pytest.mark.parametrize("dtype", DTYPES) +@pytest.mark.parametrize( + "shape", SHAPES, ids=lambda s: s[0] if isinstance(s, tuple) else str(s) +) +def test_varlen_path_matches_sdpa_on_valid_rows(dtype, shape): + """Valid rows: varlen output ≈ SDPA output within FA tolerance.""" + _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape + torch.manual_seed(0) + s = s_txt + s_img + softmax_scale = head_dim**-0.5 + mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) + q = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + k = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + v = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + + out_sdpa = _sdpa_with_key_mask(q, k, v, mask, softmax_scale) + out_varlen = _varlen_path(q, k, v, mask, softmax_scale) + + valid = mask[..., None, None].expand_as(out_sdpa) + rtol = 1e-2 if dtype == torch.bfloat16 else 5e-3 + atol = 5e-2 if dtype == torch.bfloat16 else 1e-2 + torch.testing.assert_close( + out_sdpa[valid], + out_varlen[valid], + rtol=rtol, + atol=atol, + ) + + +@pytest.mark.parametrize("dtype", DTYPES) +@pytest.mark.parametrize( + "shape", SHAPES, ids=lambda s: s[0] if isinstance(s, tuple) else str(s) +) +def test_varlen_path_zeros_masked_rows(dtype, shape): + """Masked rows: varlen path produces exact zeros (documented contract).""" + _, bs, s_txt, s_img, num_heads, head_dim, valid_txt_lens = shape + torch.manual_seed(1) + s = s_txt + s_img + softmax_scale = head_dim**-0.5 + mask = _build_mask(bs, s_txt, s_img, valid_txt_lens) + q = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + k = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + v = torch.randn(bs, s, num_heads, head_dim, dtype=dtype, device=DEVICE) + + out_varlen = _varlen_path(q, k, v, mask, softmax_scale) + + invalid = ~mask + if invalid.any(): + assert (out_varlen[invalid] == 0).all(), "masked rows must be zero-filled" + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v", "-s"])) diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/__init__.py b/python/sglang/multimodal_gen/runtime/layers/attention/__init__.py index aacb7412b..def73426c 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/__init__.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/__init__.py @@ -12,6 +12,7 @@ from sglang.multimodal_gen.runtime.layers.attention.layer import ( UlyssesAttention, UlyssesAttention_VSA, USPAttention, + build_varlen_mask_meta, ) from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend from sglang.multimodal_gen.runtime.layers.attention.turbo_layer import MinimalA2AAttnOp @@ -27,4 +28,5 @@ __all__ = [ "AttentionMetadataBuilder", # "AttentionState", "get_attn_backend", + "build_varlen_mask_meta", ] diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index 5eb38bb3f..cc4515756 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -1,6 +1,7 @@ # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # SPDX-License-Identifier: Apache-2.0 +import os from contextlib import nullcontext from typing import Type @@ -8,6 +9,12 @@ import torch import torch.nn as nn from torch.nn.attention import SDPBackend, sdpa_kernel +from sglang.jit_kernel.diffusion.triton.varlen_pack_pad import ( + build_inv_indices, + fused_pack_qkv, + fused_scatter_to_padded, +) +from sglang.jit_kernel.flash_attention import flash_attn_varlen_func from sglang.multimodal_gen.runtime.distributed.communication_op import ( sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D, @@ -20,6 +27,9 @@ from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_sp_world_size, get_ulysses_parallel_world_size, ) +from sglang.multimodal_gen.runtime.layers.attention.backends import ( + flash_attn as _fa_backend, +) from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend import ( AttentionImpl, wrap_attention_impl_forward, @@ -44,6 +54,38 @@ _PYTORCH_DEFAULT_CUDA_SDP_BACKENDS = [ SDPBackend.MATH, ] +# Set ``SGLANG_VARLEN_FA=0`` to disable the varlen FA fast path in +# USPAttention masked branch and fall back to SDPA. +_VARLEN_FA_ENABLED = os.environ.get("SGLANG_VARLEN_FA", "1") != "0" + + +def build_varlen_mask_meta( + key_mask: torch.Tensor, +) -> dict: + """Build varlen FA metadata from a ``[B, S]`` key mask. + + Returns ``cu_seqlens``, ``indices``, ``inv_indices``, ``max_seqlen``. + Passing the result via ``joint_attention_kwargs`` opts the caller into + ``USPAttention``'s varlen FA fast path, which zero-fills masked query + rows on output — only use when those rows are dropped or ignored + downstream. + """ + assert key_mask.dim() == 2, "key_mask must be [B, S]" + bs, seq = key_mask.shape + bool_mask = key_mask.to(dtype=torch.bool) + valid_lens = bool_mask.sum(dim=1, dtype=torch.int32) + indices = bool_mask.reshape(-1).nonzero(as_tuple=False).flatten() + cu_seqlens = torch.zeros(bs + 1, dtype=torch.int32, device=key_mask.device) + cu_seqlens[1:] = torch.cumsum(valid_lens, dim=0) + inv_indices = build_inv_indices(indices, bs * seq) + + return { + "cu_seqlens": cu_seqlens, + "indices": indices, + "inv_indices": inv_indices, + "max_seqlen": seq, # upper bound; FA varlen uses cu_seqlens for actual ranges + } + class UlyssesAttention(nn.Module): """Ulysses-style SequenceParallelism attention layer.""" @@ -420,6 +462,7 @@ class USPAttention(nn.Module): num_replicated_suffix: int = 0, num_replicated_kv_prefix: int = 0, skip_sequence_parallel_override: bool = False, + attn_mask_meta: dict | None = None, ) -> torch.Tensor: """ Forward pass for USPAttention. @@ -439,6 +482,10 @@ class USPAttention(nn.Module): conditioning prefix (e.g. cached text K/V) followed by a sequence-sharded suffix (image tokens). Q has no replicated portion and is fully sequence-sharded. + attn_mask_meta: optional varlen metadata from + ``build_varlen_mask_meta(attn_mask)``. Supplying this opts + into the varlen FA fast path, in which masked query rows + are zero-filled on output (differs from SDPA semantics). Note: Replicated tensors are not supported in this implementation. When skip_sequence_parallel=True (set at construction time), all SP @@ -473,6 +520,52 @@ class USPAttention(nn.Module): sp_world_size = get_sequence_parallel_world_size() if effective_skip_sp or sp_world_size == 1: + # Varlen FA fast path: SDPA with a non-None mask falls back + # to cutlassF. Meta-gated to opt in callers that drop masked + # query rows downstream (zero-filled on output, differs from + # SDPA semantics). Without meta, fall through to SDPA. + if ( + _VARLEN_FA_ENABLED + and attn_mask_meta is not None + and self.backend == AttentionBackendEnum.FA + and attn_mask.dim() == 2 + and attn_mask.dtype + in (torch.bool, torch.uint8, torch.int32, torch.int64) + 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] + ): + bs, seq = q.shape[0], q.shape[1] + indices = attn_mask_meta["indices"] + cu_seqlens = attn_mask_meta["cu_seqlens"] + max_seqlen = attn_mask_meta["max_seqlen"] + inv_indices = attn_mask_meta["inv_indices"] + # Guard against a caller passing meta from a different + # mask shape (silent corruption otherwise). + assert ( + inv_indices.shape[0] == bs * seq + ), "attn_mask_meta shape does not match attn_mask" + # All-False mask: FA varlen rejects zero-length input. + # Fall through to SDPA which handles it via broadcast. + # (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) + out_unpad = flash_attn_varlen_func( + q=q_unpad, + k=k_unpad, + v=v_unpad, + cu_seqlens_q=cu_seqlens, + cu_seqlens_k=cu_seqlens, + max_seqlen_q=max_seqlen, + max_seqlen_k=max_seqlen, + softmax_scale=self.softmax_scale, + causal=False, + ver=_fa_backend.fa_ver, + ) + return fused_scatter_to_padded(out_unpad, inv_indices, bs, seq) + q_ = q.transpose(1, 2) k_ = k.transpose(1, 2) v_ = v.transpose(1, 2) 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 694d7ea95..7bb18e9c0 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -19,7 +19,10 @@ from sglang.multimodal_gen.runtime.distributed import get_local_torch_device from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_sp_world_size, ) -from sglang.multimodal_gen.runtime.layers.attention import USPAttention +from sglang.multimodal_gen.runtime.layers.attention import ( + USPAttention, + build_varlen_mask_meta, +) from sglang.multimodal_gen.runtime.layers.elementwise import MulAdd from sglang.multimodal_gen.runtime.layers.fused_scale_shift_gate import ( FusedLayerNormScaleShiftGateSelect01, @@ -653,6 +656,9 @@ class QwenImageCrossAttention(nn.Module): encoder_hidden_states_mask = cross_attention_kwargs.get( "encoder_hidden_states_mask" ) + # Varlen metadata precomputed in QwenImageTransformer2DModel.forward, + # paired with the same ``attn_mask`` for the USPAttention FA fast path. + attn_mask_meta = cross_attention_kwargs.get("attn_mask_meta") ( img_query, @@ -733,6 +739,7 @@ class QwenImageCrossAttention(nn.Module): joint_key, joint_value, attn_mask=attn_mask, + attn_mask_meta=attn_mask_meta, num_replicated_prefix=seq_len_txt, ) @@ -1351,8 +1358,12 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): dtype=torch.bool, device=hidden_states.device, ) - block_attention_kwargs["attn_mask"] = torch.cat( - [encoder_hidden_states_mask, image_mask], dim=1 + joint_mask = torch.cat([encoder_hidden_states_mask, image_mask], dim=1) + block_attention_kwargs["attn_mask"] = joint_mask + # Precompute varlen metadata once per request so every block reuses + # the same cu_seqlens / indices instead of rebuilding. + block_attention_kwargs["attn_mask_meta"] = build_varlen_mask_meta( + joint_mask ) temb = self.time_text_embed(timestep, hidden_states, additional_t_cond)