[diffusion] optimize: shard qwen text embed in sp (#29147)

This commit is contained in:
Aleksi Vesanto
2026-06-26 21:44:45 +08:00
committed by GitHub
parent b73e57210a
commit b91348071e
@@ -17,6 +17,8 @@ from diffusers.models.normalization import AdaLayerNormContinuous
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
get_ring_parallel_world_size,
get_sp_parallel_rank,
get_sp_world_size,
)
from sglang.multimodal_gen.runtime.layers.attention import (
@@ -73,6 +75,99 @@ def _local_seq_len(seq_len: int, sp_world_size: int) -> int:
return padded_len // sp_world_size
def _shard_text_for_sp(
encoder_hidden_states: torch.Tensor,
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]],
) -> Tuple[torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor]]]:
"""Shard the replicated text stream evenly across SP ranks.
The image latents are already sharded by the pipeline while the text stream
is replicated. This splits the text embeddings (and their RoPE cache) so each
rank keeps ``1/sp_size`` of the text tokens, making the joint attention fully
sequence-sharded (``num_replicated_prefix=0``). Callers must ensure the text
length divides evenly across SP ranks.
"""
sp_size = get_sp_world_size()
if sp_size == 1:
return encoder_hidden_states, freqs_cis
sp_rank = get_sp_parallel_rank()
encoder_hidden_states = torch.chunk(encoder_hidden_states, sp_size, dim=1)[sp_rank]
if freqs_cis is not None:
img_cache, txt_cache = freqs_cis
txt_cache = torch.chunk(txt_cache, sp_size, dim=0)[sp_rank]
freqs_cis = (img_cache, txt_cache)
return encoder_hidden_states, freqs_cis
def _pad_shard_text_for_sp_varlen(
encoder_hidden_states: torch.Tensor,
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]],
image_seq_len: int,
) -> Tuple[
torch.Tensor,
Optional[Tuple[torch.Tensor, torch.Tensor]],
torch.Tensor,
Dict[str, int],
]:
"""Right-pad a non-divisible replicated text stream to a multiple of the SP
world size and shard it evenly across ranks, so the joint ``[text, image]``
sequence is fully sequence-parallel.
The pad tokens occupy a single contiguous block at the tail of the last
rank's text chunk. The returned ``attn_mask`` (joint ``[text, image]``
validity mask) and ``attn_mask_meta`` (``gap_start`` / ``gap_end``) describe
that block so ``USPAttention`` excludes it from attention via the varlen
kernel.
Callers must ensure the text length is NOT divisible by the SP world size;
the evenly-divisible case is a plain ``_shard_text_for_sp`` with no mask.
Returns ``(encoder_hidden_states, freqs_cis, attn_mask, attn_mask_meta)``.
"""
sp_size = get_sp_world_size()
t_real = encoder_hidden_states.shape[1]
num_pad = sp_size - t_real % sp_size
encoder_hidden_states = F.pad(encoder_hidden_states, (0, 0, 0, num_pad))
if freqs_cis is not None:
img_cache, txt_cache = freqs_cis
txt_cache = F.pad(txt_cache, (0, 0, 0, num_pad))
freqs_cis = (img_cache, txt_cache)
local_txt = (t_real + num_pad) // sp_size
encoder_hidden_states, freqs_cis = _shard_text_for_sp(
encoder_hidden_states, freqs_cis
)
sp_rank = get_sp_parallel_rank()
txt_start = sp_rank * local_txt
valid_txt = min(local_txt, max(t_real - txt_start, 0))
text_mask = torch.zeros(
encoder_hidden_states.shape[0],
local_txt,
dtype=torch.bool,
device=encoder_hidden_states.device,
)
text_mask[:, :valid_txt] = True
image_mask = torch.ones(
encoder_hidden_states.shape[0],
image_seq_len,
dtype=torch.bool,
device=encoder_hidden_states.device,
)
joint_mask = torch.cat([text_mask, image_mask], dim=1)
# Gathered joint layout is rank-major [txt_0, img_0, ..., txt_{sp-1},
# img_{sp-1}]; the pad block is the tail of the last rank's text chunk.
gap_meta = {
"gap_start": (sp_size - 1) * (local_txt + image_seq_len) + local_txt - num_pad,
"gap_end": (sp_size - 1) * (local_txt + image_seq_len) + local_txt,
}
return encoder_hidden_states, freqs_cis, joint_mask, gap_meta
def _get_qkv_projections(
attn: "QwenImageCrossAttention", hidden_states, encoder_hidden_states=None
):
@@ -659,6 +754,9 @@ class QwenImageCrossAttention(nn.Module):
# 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")
# When the text stream is sharded across SP ranks the joint sequence is
# fully sequence-parallel, so no leading tokens are replicated.
sp_text_sharded = cross_attention_kwargs.get("sp_text_sharded", False)
(
img_query,
@@ -740,7 +838,7 @@ class QwenImageCrossAttention(nn.Module):
joint_value,
attn_mask=attn_mask,
attn_mask_meta=attn_mask_meta,
num_replicated_prefix=seq_len_txt,
num_replicated_prefix=0 if sp_text_sharded else seq_len_txt,
)
# Reshape back
@@ -1348,6 +1446,8 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
encoder_hidden_states, _ = self.txt_in(encoder_hidden_states)
block_attention_kwargs = attention_kwargs.copy() if attention_kwargs else {}
sp_text_sharded = False
sp_size = get_sp_world_size()
if encoder_hidden_states_mask is not None:
encoder_hidden_states_mask = encoder_hidden_states_mask.to(
device=hidden_states.device, dtype=torch.bool
@@ -1365,6 +1465,30 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
block_attention_kwargs["attn_mask_meta"] = build_varlen_mask_meta(
joint_mask
)
elif sp_size > 1 and encoder_hidden_states.shape[1] % sp_size == 0:
# Text divides evenly across SP ranks: plain even shard, no mask.
encoder_hidden_states, freqs_cis = _shard_text_for_sp(
encoder_hidden_states, freqs_cis
)
sp_text_sharded = True
elif sp_size > 1 and get_ring_parallel_world_size() == 1:
# Text does not divide evenly: pad to an SP multiple and shard, with a
# pad-gap mask so USPAttention excludes the padding via the varlen
# kernel. The varlen masked path does not support ring parallelism, so
# uneven text under ring>1 instead falls through to the replicated
# path below.
(
encoder_hidden_states,
freqs_cis,
pad_mask,
pad_meta,
) = _pad_shard_text_for_sp_varlen(
encoder_hidden_states, freqs_cis, hidden_states.shape[1]
)
sp_text_sharded = True
block_attention_kwargs["attn_mask"] = pad_mask
block_attention_kwargs["attn_mask_meta"] = pad_meta
block_attention_kwargs["sp_text_sharded"] = sp_text_sharded
temb = self.time_text_embed(timestep, hidden_states, additional_t_cond)