[diffusion] perf: build qwen's masked varlen metadata host-side (#33954)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-08-07 17:55:15 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent 0756a1d2b0
commit 28b43bf693
2 changed files with 57 additions and 0 deletions
@@ -39,6 +39,7 @@ from sglang.multimodal_gen.runtime.layers.attention import (
DynamicVarlenMaskMeta,
USPAttention,
build_varlen_mask_meta,
build_varlen_mask_meta_from_ranges,
)
from sglang.multimodal_gen.runtime.layers.elementwise import MulAdd
from sglang.multimodal_gen.runtime.layers.fused_scale_shift_gate import (
@@ -1565,6 +1566,26 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
# once, so build varlen metadata replay-locally from the current
# static mask instead of closing over stale cu_seqlens/indices.
block_attention_kwargs["attn_mask_meta"] = DynamicVarlenMaskMeta()
elif (
txt_seq_lens is not None
and len(txt_seq_lens) == batch_size
and all(0 <= n <= encoder_hidden_states.shape[1] for n in txt_seq_lens)
):
# txt_seq_lens already carries each row's valid text prefix
# (the mask is built from it), so the varlen metadata can be
# assembled host-side; the mask-based builder costs a GPU
# nonzero plus a device sync on every denoising step.
txt_len = encoder_hidden_states.shape[1]
block_attention_kwargs["attn_mask_meta"] = (
build_varlen_mask_meta_from_ranges(
[
[(0, int(n)), (txt_len, txt_len + image_seq_len)]
for n in txt_seq_lens
],
max_seqlen=txt_len + image_seq_len,
device=hidden_states.device,
)
)
else:
# Precompute varlen metadata once per request so every block
# reuses the same cu_seqlens / indices instead of rebuilding.
@@ -0,0 +1,36 @@
"""Host-built varlen metadata must match the mask-based (nonzero) builder."""
import unittest
import torch
from sglang.multimodal_gen.runtime.layers.attention.layer import (
build_varlen_mask_meta,
build_varlen_mask_meta_from_ranges,
)
class TestHostVarlenMetaEquivalence(unittest.TestCase):
def test_prefix_text_plus_full_image_matches_nonzero_builder(self):
txt_len, img_len = 7, 5
txt_seq_lens = [3, 7, 0]
bs = len(txt_seq_lens)
mask = torch.zeros(bs, txt_len + img_len, dtype=torch.bool)
for row, n in enumerate(txt_seq_lens):
mask[row, :n] = True
mask[row, txt_len:] = True
ref = build_varlen_mask_meta(mask)
host = build_varlen_mask_meta_from_ranges(
[[(0, n), (txt_len, txt_len + img_len)] for n in txt_seq_lens],
max_seqlen=txt_len + img_len,
device=mask.device,
)
for key in ("cu_seqlens", "indices", "inv_indices"):
torch.testing.assert_close(host[key], ref[key], rtol=0, atol=0)
self.assertEqual(host["max_seqlen"], ref["max_seqlen"])
if __name__ == "__main__":
unittest.main()