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 7bb18e9c0..cd37330c4 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -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)