[diffusion] optimize: shard qwen text embed in sp (#29147)
This commit is contained in:
@@ -17,6 +17,8 @@ from diffusers.models.normalization import AdaLayerNormContinuous
|
|||||||
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
|
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 import get_local_torch_device
|
||||||
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
from sglang.multimodal_gen.runtime.distributed.parallel_state import (
|
||||||
|
get_ring_parallel_world_size,
|
||||||
|
get_sp_parallel_rank,
|
||||||
get_sp_world_size,
|
get_sp_world_size,
|
||||||
)
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import (
|
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
|
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(
|
def _get_qkv_projections(
|
||||||
attn: "QwenImageCrossAttention", hidden_states, encoder_hidden_states=None
|
attn: "QwenImageCrossAttention", hidden_states, encoder_hidden_states=None
|
||||||
):
|
):
|
||||||
@@ -659,6 +754,9 @@ class QwenImageCrossAttention(nn.Module):
|
|||||||
# Varlen metadata precomputed in QwenImageTransformer2DModel.forward,
|
# Varlen metadata precomputed in QwenImageTransformer2DModel.forward,
|
||||||
# paired with the same ``attn_mask`` for the USPAttention FA fast path.
|
# paired with the same ``attn_mask`` for the USPAttention FA fast path.
|
||||||
attn_mask_meta = cross_attention_kwargs.get("attn_mask_meta")
|
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,
|
img_query,
|
||||||
@@ -740,7 +838,7 @@ class QwenImageCrossAttention(nn.Module):
|
|||||||
joint_value,
|
joint_value,
|
||||||
attn_mask=attn_mask,
|
attn_mask=attn_mask,
|
||||||
attn_mask_meta=attn_mask_meta,
|
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
|
# Reshape back
|
||||||
@@ -1348,6 +1446,8 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
|||||||
encoder_hidden_states, _ = self.txt_in(encoder_hidden_states)
|
encoder_hidden_states, _ = self.txt_in(encoder_hidden_states)
|
||||||
|
|
||||||
block_attention_kwargs = attention_kwargs.copy() if attention_kwargs else {}
|
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:
|
if encoder_hidden_states_mask is not None:
|
||||||
encoder_hidden_states_mask = encoder_hidden_states_mask.to(
|
encoder_hidden_states_mask = encoder_hidden_states_mask.to(
|
||||||
device=hidden_states.device, dtype=torch.bool
|
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(
|
block_attention_kwargs["attn_mask_meta"] = build_varlen_mask_meta(
|
||||||
joint_mask
|
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)
|
temb = self.time_text_embed(timestep, hidden_states, additional_t_cond)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user