diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 039e95655..a7d4fd0af 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -145,37 +145,15 @@ def pad_text_embeddings_with_mask( def shard_rotary_emb_for_sp(emb): - """ - Shard rotary embeddings [S, D] along sequence for SP. - If S is not divisible by SP degree, pad by repeating the last row. - """ - # Sequence Parallelism: slice image RoPE to local shard if enabled + """Shard rotary embeddings [S, D] along the sequence for SP; non-divisible + lengths pad by repeating the last row (position labels, never attention + K/V, so the pad value only needs to stay finite).""" try: - from sglang.multimodal_gen.runtime.distributed.parallel_state import ( - get_sp_parallel_rank, - get_sp_world_size, - ) + from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import shard_seq - sp_world_size = get_sp_world_size() + return shard_seq(emb, dim=0, pad_mode="repeat_last")[0] except Exception: - sp_world_size = 1 - seq_len = emb.shape[0] - if seq_len % sp_world_size != 0: - pad_len = sp_world_size - (seq_len % sp_world_size) - pad = emb[-1:].repeat(pad_len, 1) - emb = torch.cat([emb, pad], dim=0) - if sp_world_size > 1: - try: - rank = get_sp_parallel_rank() - except Exception: - rank = 0 - seq_len = emb.shape[0] - local_len = seq_len // sp_world_size - start = rank * local_len - end = start + local_len - emb = emb[start:end] - return emb - else: + # Distributed state not initialized (single-process utilities). return emb @@ -543,18 +521,15 @@ class PipelineConfig: return latents, False time_dim = latents.shape[2] - # Pad to next multiple of SP degree if needed + # Zero-padding a non-divisible time dim would enter self-attention + # unmasked (video models pass no attn_mask) and corrupt real tokens; + # keep such shapes unsharded until models consume the sp_shard meta. if time_dim > 0 and time_dim % sp_world_size != 0: - logger.debug( - "Padding latents to next multiple of SP degree, performance is sub-optimal" + logger.warning_once( + f"Latent time dim {time_dim} is not divisible by SP degree " + f"{sp_world_size}; skipping sequence shard for correctness." ) - pad_len = sp_world_size - (time_dim % sp_world_size) - pad = torch.zeros( - (*latents.shape[:2], pad_len, *latents.shape[3:]), - dtype=latents.dtype, - device=latents.device, - ) - latents = torch.cat([latents, pad], dim=2) + return latents, False assert latents.shape[2] % sp_world_size == 0 sharded_tensor = rearrange( diff --git a/python/sglang/multimodal_gen/runtime/distributed/sp_shard_utils.py b/python/sglang/multimodal_gen/runtime/distributed/sp_shard_utils.py new file mode 100644 index 000000000..b589798ee --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/distributed/sp_shard_utils.py @@ -0,0 +1,225 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Unified sequence-parallel shard / pad / gather helpers. + +Layout invariant: padding always sits at the end of the LAST rank's local +chunk, so the ulysses-gathered sequence carries one contiguous pad block at its +global tail. `tail_attn_meta` then lets attention skip that block for free +(the pad becomes its own varlen segment - no repacking, no mask compute). +""" + +from __future__ import annotations + +import os +from dataclasses import dataclass + +import torch +import torch.nn.functional as F + +from sglang.multimodal_gen.runtime.distributed.communication_op import ( + sequence_model_parallel_all_gather, +) +from sglang.multimodal_gen.runtime.distributed.parallel_state import ( + get_ring_parallel_world_size, + get_sp_parallel_rank, + get_sp_world_size, +) + +# Text shorter than this stays replicated instead of SP-sharded (see +# plan_text_strategy). 0 = always shard when legal; H100 bench showed sharding +# wins from trivial lengths on, so the knob exists only as an escape hatch. +_TEXT_SHARD_MIN = int(os.environ.get("SGLANG_SP_TEXT_SHARD_MIN", "0")) + + +@dataclass(frozen=True) +class SpShard: + """Facts of one tail-padded even shard, shared by tensors of that stream.""" + + orig_len: int # real tokens (global) + local_len: int # per-rank chunk length (equal on every rank) + num_pad: int # pad tokens, all at the last rank's local tail + sp_size: int + sp_rank: int + + @property + def local_pad(self) -> int: + """Pad rows inside THIS rank's chunk (tail rows of the last rank).""" + return self.num_pad if self.sp_rank == self.sp_size - 1 else 0 + + @property + def local_real_len(self) -> int: + return self.local_len - self.local_pad + + +def build_shard_plan(seq_len: int) -> SpShard: + """Shard math only; tensors are sliced separately via `shard_like`.""" + sp_size = get_sp_world_size() + if sp_size <= 1: + return SpShard(seq_len, seq_len, 0, 1, 0) + local_len = (seq_len + sp_size - 1) // sp_size + return SpShard( + orig_len=seq_len, + local_len=local_len, + num_pad=local_len * sp_size - seq_len, + sp_size=sp_size, + sp_rank=get_sp_parallel_rank(), + ) + + +def shard_like( + x: torch.Tensor, shard: SpShard, dim: int = 1, pad_mode: str = "zeros" +) -> torch.Tensor: + """Apply a planned shard to one tensor (RoPE caches use the same plan as + hidden states so their chunks stay aligned).""" + if shard.sp_size <= 1: + return x + if shard.num_pad > 0: + if pad_mode == "repeat_last": + pad = x.narrow(dim, x.shape[dim] - 1, 1) + pad = pad.expand( + *[shard.num_pad if i == dim else -1 for i in range(x.dim())] + ) + x = torch.cat([x, pad], dim=dim) + else: + # F.pad pads dims last-to-first: (left, right) pairs from dim -1. + pads = [0, 0] * (x.dim() - 1 - dim) + [0, shard.num_pad] + x = F.pad(x, pads) + return x.narrow(dim, shard.sp_rank * shard.local_len, shard.local_len) + + +def shard_seq( + x: torch.Tensor, dim: int = 1, pad_mode: str = "zeros" +) -> tuple[torch.Tensor, SpShard]: + """ + mode: + zeroes: pad with zeroes at tail + repeat_last: repeat the last token, only for rotary embedding + """ + shard = build_shard_plan(x.shape[dim]) + return shard_like(x, shard, dim=dim, pad_mode=pad_mode), shard + + +def gather_seq(local: torch.Tensor, orig_len: int, dim: int = 1) -> torch.Tensor: + """All-gather an SP-sharded sequence and trim the tail padding""" + if get_sp_world_size() <= 1: + return local + full = sequence_model_parallel_all_gather(local.contiguous(), dim=dim) + if full.shape[dim] > orig_len: + full = full.narrow(dim, 0, orig_len) + return full + + +def shard_seq_prefix( + x: torch.Tensor, prefix_len: int, shard: SpShard, dim: int = 0 +) -> torch.Tensor: + """Shard only the leading ``prefix_len`` rows (e.g. the text segment of a + joint RoPE cache) with an existing plan; the remainder is kept as-is.""" + rest = x.shape[dim] - prefix_len + return torch.cat( + [ + shard_like(x.narrow(dim, 0, prefix_len), shard, dim=dim), + x.narrow(dim, prefix_len, rest), + ], + dim=dim, + ) + + +def join_seqs( + prefix: torch.Tensor, body: torch.Tensor, local_pad: int, dim: int = 1 +) -> torch.Tensor: + """Concatenate local sharded ``[prefix (txt tokens, padding tokens), body (img tokens)]`` for joint attention, while relocating the + prefix's ``local_pad`` tail rows behind the body. + + Why leave the padding at tail: the shard pads the *text* chunk, but the local joint layout is + [text, image]. + + In naive implementation, after the ulysses all-to-all, that pad would sit mid-sequence (of last rank) + ([... txt_last, PAD, img_last]), which required further mem copy (for the padding tokens), inefficient in this case + + With the pad relocated behind the image, the padding forms one global-tail block that the zero-copy varlen + path (tail_attn_meta, implemented in USPAttention.forward) skips for free + """ + if local_pad > 0: + real = prefix.shape[dim] - local_pad + return torch.cat( + [ + # txt tokens + prefix.narrow(dim, 0, real), + body, + # leave the padding at global-tail + prefix.narrow(dim, real, local_pad), + ], + dim=dim, + ) + return torch.cat([prefix, body], dim=dim) + + +def split_seqs( + joint: torch.Tensor, prefix_len: int, local_pad: int, dim: int = 1 +) -> tuple[torch.Tensor, torch.Tensor]: + """Inverse of ``join_seqs``: recover ``(prefix, body)`` from the joint output, with the pad rows rejoining the prefix tail so the residual text + stream keeps its per-rank shape. + + ([... txt_last, PAD, img_last]) -> prefix (txt + pad), body (img) + """ + total = joint.shape[dim] + if local_pad > 0: + real = prefix_len - local_pad + body_end = total - local_pad + prefix = torch.cat( + [joint.narrow(dim, 0, real), joint.narrow(dim, body_end, local_pad)], + dim=dim, + ) + return prefix, joint.narrow(dim, real, body_end - real) + return ( + joint.narrow(dim, 0, prefix_len), + joint.narrow(dim, prefix_len, total - prefix_len), + ) + + +def should_shard_text(txt_len: int) -> bool: + """True when the joint-attention text stream should be SP-sharded here + (see plan_text_strategy for the policy).""" + return get_sp_world_size() > 1 and plan_text_strategy(txt_len) == "shard" + + +def tail_attn_meta( + shard: SpShard, + batch_size: int, + device: torch.device, + image_seq_len: int = 0, +) -> dict | None: + """Per-request attention meta for a tail-padded shard: `cu_seqlens_tail` + splits each batch row into [valid | pad] varlen segments over the gathered + layout, so USPAttention runs varlen FA on the padded q/k/v with zero + repacking. Built once per request, reused by every block.""" + if shard.sp_size <= 1 or shard.num_pad == 0: + return None + seq = shard.sp_size * (shard.local_len + image_seq_len) + valid = seq - shard.num_pad + row = torch.tensor([valid, shard.num_pad], dtype=torch.int32, device=device) + seglens = row.repeat(batch_size) + cu_seqlens = torch.zeros(2 * batch_size + 1, dtype=torch.int32, device=device) + cu_seqlens[1:] = torch.cumsum(seglens, dim=0) + return { + "pad_start": valid, + "pad_end": seq, + "local_pad": shard.local_pad, + "cu_seqlens_tail": cu_seqlens, + "max_seqlen_tail": max(valid, shard.num_pad), + } + + +def plan_text_strategy(txt_len: int) -> str: + """Choose "shard" or "replicate" for the joint-attention text stream. + + Prefer "shard" by default. for small sequence (shorter than SGLANG_SP_TEXT_SHARD_MIN), choose "replicate" for better performance + + """ + sp_size = get_sp_world_size() + if sp_size <= 1: + return "replicate" + if txt_len % sp_size != 0 and get_ring_parallel_world_size() > 1: + return "replicate" + if txt_len < _TEXT_SHARD_MIN: + return "replicate" + return "shard" diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index 0039b4935..24e82377a 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -533,7 +533,21 @@ class USPAttention(nn.Module): effective_skip_sp = ( self.skip_sequence_parallel or skip_sequence_parallel_override ) - if attn_mask is not None: + # Tail-pad meta alone (sp_shard.tail_attn_meta; mask derivable from the + # pad span) also opts into the masked SP branch. gap_* = legacy alias. + meta_pad_start = meta_pad_end = None + if attn_mask_meta is not None: + meta_pad_start = attn_mask_meta.get( + "pad_start", attn_mask_meta.get("gap_start") + ) + meta_pad_end = attn_mask_meta.get("pad_end", attn_mask_meta.get("gap_end")) + meta_only_pad = ( + attn_mask is None + and meta_pad_start is not None + and not effective_skip_sp + and get_sequence_parallel_world_size() > 1 + ) + if attn_mask is not None or meta_only_pad: def _prepare_sdpa_mask( mask: torch.Tensor, *, dtype: torch.dtype, device: torch.device @@ -626,7 +640,7 @@ class USPAttention(nn.Module): raise NotImplementedError( "USPAttention masked path does not support ring parallelism yet." ) - if attn_mask.dim() != 2: + if attn_mask is not None and attn_mask.dim() != 2: raise NotImplementedError( "USPAttention masked SP path currently expects a [B, S_local] key mask." ) @@ -637,26 +651,46 @@ class USPAttention(nn.Module): k = _usp_input_all_to_all(k, head_dim=2) v = _usp_input_all_to_all(v, head_dim=2) - gap_start = None - gap_end = None - if attn_mask_meta is not None: - gap_start = attn_mask_meta.get("gap_start") - gap_end = attn_mask_meta.get("gap_end") if ( _VARLEN_FA_ENABLED and self.backend == AttentionBackendEnum.FA - and gap_start is not None - and gap_end is not None - and gap_end > gap_start + and meta_pad_start is not None + and meta_pad_end is not None + and meta_pad_end > meta_pad_start and q.device.type == "cuda" and q.dtype in (torch.float16, torch.bfloat16) ): bs, seq = q.shape[0], q.shape[1] - assert 0 <= gap_start < gap_end <= seq - valid_seq = seq - (gap_end - gap_start) - q_dense = torch.cat([q[:, :gap_start], q[:, gap_end:]], dim=1) - k_dense = torch.cat([k[:, :gap_start], k[:, gap_end:]], dim=1) - v_dense = torch.cat([v[:, :gap_start], v[:, gap_end:]], dim=1) + assert 0 <= meta_pad_start < meta_pad_end <= seq + cu_tail = attn_mask_meta.get("cu_seqlens_tail") + if cu_tail is not None and meta_pad_end == seq: + # Zero-copy tail path: run varlen FA straight over the + # padded layout, each row split into [valid | pad] segments + # (contiguous reshapes only, no repacking). + assert ( + cu_tail.numel() == 2 * bs + 1 + ), "cu_seqlens_tail does not match the batch size" + out = flash_attn_varlen_func( + q=q.reshape(bs * seq, *q.shape[2:]), + k=k.reshape(bs * seq, *k.shape[2:]), + v=v.reshape(bs * seq, *v.shape[2:]), + cu_seqlens_q=cu_tail, + cu_seqlens_k=cu_tail, + max_seqlen_q=attn_mask_meta["max_seqlen_tail"], + max_seqlen_k=attn_mask_meta["max_seqlen_tail"], + softmax_scale=self.softmax_scale, + causal=False, + ver=_fa_backend.fa_ver, + ).reshape(bs, seq, *q.shape[2:]) + # Match the packed paths: masked query rows read as zeros. + out[:, meta_pad_start:].zero_() + if sp_size > 1: + out = _usp_output_all_to_all(out, head_dim=2) + return out + valid_seq = seq - (meta_pad_end - meta_pad_start) + q_dense = torch.cat([q[:, :meta_pad_start], q[:, meta_pad_end:]], dim=1) + k_dense = torch.cat([k[:, :meta_pad_start], k[:, meta_pad_end:]], dim=1) + v_dense = torch.cat([v[:, :meta_pad_start], v[:, meta_pad_end:]], dim=1) cu_seqlens = torch.arange( 0, (bs + 1) * valid_seq, @@ -677,10 +711,17 @@ class USPAttention(nn.Module): ver=_fa_backend.fa_ver, ).reshape(bs, valid_seq, *q.shape[2:]) gap_out = out_dense.new_zeros( - bs, gap_end - gap_start, out_dense.shape[2], out_dense.shape[3] + bs, + meta_pad_end - meta_pad_start, + out_dense.shape[2], + out_dense.shape[3], ) out = torch.cat( - [out_dense[:, :gap_start], gap_out, out_dense[:, gap_start:]], + [ + out_dense[:, :meta_pad_start], + gap_out, + out_dense[:, meta_pad_start:], + ], dim=1, ) if sp_size > 1: @@ -691,9 +732,17 @@ class USPAttention(nn.Module): # attn_mask is inconsistent across SP ranks (None on some, Tensor on # others), which causes all_gather participant mismatch. Upstream # mask builders must ensure all ranks produce the same mask type. - gathered_mask = sequence_model_parallel_all_gather( - attn_mask.contiguous(), dim=1 - ) + if attn_mask is None: + # Meta-only tail-pad caller on a non-FA fallback: the gathered + # mask is fully determined by the pad span, no collective needed. + gathered_mask = torch.ones( + q.shape[0], q.shape[1], dtype=torch.bool, device=q.device + ) + gathered_mask[:, meta_pad_start:meta_pad_end] = False + else: + gathered_mask = sequence_model_parallel_all_gather( + attn_mask.contiguous(), dim=1 + ) if ( _VARLEN_FA_ENABLED and self.backend == AttentionBackendEnum.FA diff --git a/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py b/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py index 082f11ed8..5881220bf 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/ernie_image.py @@ -134,10 +134,14 @@ class ErnieImageSelfAttention(nn.Module): self.norm_q = RMSNorm(head_dim, eps=eps) self.norm_k = RMSNorm(head_dim, eps=eps) + # The joint [image, text] stream is fully replicated, so the ulysses + # all-to-all would wrongly treat it as sharded and duplicate it. Skip + # SP until the stream is sharded (sp_shard + num_replicated_suffix). self.attn = USPAttention( num_heads=self.num_local_heads, head_size=head_dim, prefix=f"{prefix}.attn", + skip_sequence_parallel=True, ) def forward( diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux.py b/python/sglang/multimodal_gen/runtime/models/dits/flux.py index b0a1c85dd..793cf0338 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux.py @@ -30,10 +30,17 @@ from torch.nn import LayerNorm as LayerNorm from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.distributed import ( divide, - get_sp_parallel_rank, - get_sp_world_size, get_tp_world_size, ) +from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import ( + build_shard_plan, + join_seqs, + shard_like, + shard_seq_prefix, + should_shard_text, + split_seqs, + tail_attn_meta, +) from sglang.multimodal_gen.runtime.layers.attention import USPAttention from sglang.multimodal_gen.runtime.layers.layernorm import ( RMSNorm, @@ -69,99 +76,6 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name -def _shard_text_for_sp( - encoder_hidden_states: torch.Tensor, - freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]], - image_seq_len: int, - num_txt_tokens: int, -) -> Tuple[ - torch.Tensor, - Optional[Tuple[torch.Tensor, torch.Tensor]], - int, - Optional[torch.Tensor], - Optional[Dict[str, int]], -]: - sp_size = get_sp_world_size() - num_replicated_prefix = num_txt_tokens - if sp_size == 1: - return encoder_hidden_states, freqs_cis, num_replicated_prefix, None, None - - sp_rank = get_sp_parallel_rank() - local_txt_tokens = (num_txt_tokens + sp_size - 1) // sp_size - padded_txt_tokens = local_txt_tokens * sp_size - num_pad_tokens = padded_txt_tokens - num_txt_tokens - - if num_pad_tokens > 0: - pad_hidden_states = encoder_hidden_states.new_zeros( - encoder_hidden_states.shape[0], - num_pad_tokens, - encoder_hidden_states.shape[2], - ) - encoder_hidden_states = torch.cat( - [encoder_hidden_states, pad_hidden_states], dim=1 - ) - - encoder_hidden_states = torch.chunk(encoder_hidden_states, sp_size, dim=1)[sp_rank] - if freqs_cis is not None: - cos, sin = freqs_cis - txt_cos = cos[:num_txt_tokens] - txt_sin = sin[:num_txt_tokens] - if num_pad_tokens > 0: - pad_cos = txt_cos.new_ones(num_pad_tokens, txt_cos.shape[1]) - pad_sin = txt_sin.new_zeros(num_pad_tokens, txt_sin.shape[1]) - txt_cos = torch.cat([txt_cos, pad_cos], dim=0) - txt_sin = torch.cat([txt_sin, pad_sin], dim=0) - freqs_cis = ( - torch.cat( - [ - torch.chunk(txt_cos, sp_size, dim=0)[sp_rank], - cos[num_txt_tokens:], - ], - dim=0, - ), - torch.cat( - [ - torch.chunk(txt_sin, sp_size, dim=0)[sp_rank], - sin[num_txt_tokens:], - ], - dim=0, - ), - ) - - num_replicated_prefix = 0 - if num_pad_tokens == 0: - return encoder_hidden_states, freqs_cis, num_replicated_prefix, None, None - - txt_start = sp_rank * local_txt_tokens - valid_txt_tokens = min(local_txt_tokens, max(num_txt_tokens - txt_start, 0)) - text_mask = torch.zeros( - encoder_hidden_states.shape[0], - local_txt_tokens, - dtype=torch.bool, - device=encoder_hidden_states.device, - ) - text_mask[:, :valid_txt_tokens] = True - image_mask = torch.ones( - encoder_hidden_states.shape[0], - image_seq_len, - dtype=torch.bool, - device=encoder_hidden_states.device, - ) - return ( - encoder_hidden_states, - freqs_cis, - num_replicated_prefix, - torch.cat([text_mask, image_mask], dim=1), - { - "gap_start": (sp_size - 1) * (local_txt_tokens + image_seq_len) - + local_txt_tokens - - num_pad_tokens, - "gap_end": (sp_size - 1) * (local_txt_tokens + image_seq_len) - + local_txt_tokens, - }, - ) - - try: from nunchaku.models.attention import NunchakuFeedForward # type: ignore[import] from nunchaku.models.normalization import ( # type: ignore[import] @@ -600,9 +514,12 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin): allow_inplace=True, ) - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) + # join_seqs relocates any SP text tail-pad behind the image (see + # sp_shard.join_seqs for why). + sp_txt_pad = (attn_mask_meta or {}).get("local_pad", 0) + query = join_seqs(encoder_query, query, sp_txt_pad) + key = join_seqs(encoder_key, key, sp_txt_pad) + value = join_seqs(encoder_value, value, sp_txt_pad) else: query, key = apply_qk_norm_with_optional_rope( q=query, @@ -627,12 +544,8 @@ class FluxAttention(torch.nn.Module, AttentionModuleMixin): x = x.to(query.dtype) if encoder_hidden_states is not None: - encoder_hidden_states, x = x.split_with_sizes( - [ - encoder_hidden_states.shape[1], - x.shape[1] - encoder_hidden_states.shape[1], - ], - dim=1, + encoder_hidden_states, x = split_seqs( + x, encoder_hidden_states.shape[1], sp_txt_pad ) if not self.pre_only: x, _ = self.to_out[0](x) @@ -775,11 +688,16 @@ class FluxSingleTransformerBlock(nn.Module): num_replicated_prefix: int = 0, ) -> Tuple[torch.Tensor, torch.Tensor]: text_seq_len = encoder_hidden_states.shape[1] - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) + joint_attention_kwargs = joint_attention_kwargs or {} + # join_seqs relocates any SP text tail-pad behind the image; the caller + # hands single blocks a RoPE cache reordered the same way. + sp_txt_pad = (joint_attention_kwargs.get("attn_mask_meta") or {}).get( + "local_pad", 0 + ) + hidden_states = join_seqs(encoder_hidden_states, hidden_states, sp_txt_pad) residual = hidden_states norm_hidden_states, gate = self.norm(hidden_states, emb=temb) - joint_attention_kwargs = joint_attention_kwargs or {} if self.use_nunchaku_structure: if _nunchaku_fused_ops_available: @@ -824,9 +742,8 @@ class FluxSingleTransformerBlock(nn.Module): if hidden_states.dtype == torch.float16: hidden_states = hidden_states.clip(-65504, 65504) - encoder_hidden_states, hidden_states = ( - hidden_states[:, :text_seq_len], - hidden_states[:, text_seq_len:], + encoder_hidden_states, hidden_states = split_seqs( + hidden_states, text_seq_len, sp_txt_pad ) return encoder_hidden_states, hidden_states @@ -1186,24 +1103,41 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): num_txt_tokens = encoder_hidden_states.shape[1] encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states) - ( - encoder_hidden_states, - freqs_cis, - num_replicated_prefix, - attn_mask, - attn_mask_meta, - ) = _shard_text_for_sp( - encoder_hidden_states, - freqs_cis, - hidden_states.shape[1], - num_txt_tokens, - ) - if attn_mask is not None: - joint_attention_kwargs = ( - joint_attention_kwargs.copy() if joint_attention_kwargs else {} + # Shard the replicated text stream across SP ranks (image latents are + # already sharded); non-divisible lengths tail-pad the last rank and the + # per-request tail meta lets attention skip the pad for free. + num_replicated_prefix = num_txt_tokens + singles_freqs_cis = freqs_cis + if should_shard_text(num_txt_tokens): + txt_shard = build_shard_plan(num_txt_tokens) + encoder_hidden_states = shard_like(encoder_hidden_states, txt_shard) + if freqs_cis is not None: + cos, sin = freqs_cis + cos = shard_seq_prefix(cos, num_txt_tokens, txt_shard) + sin = shard_seq_prefix(sin, num_txt_tokens, txt_shard) + freqs_cis = (cos, sin) + singles_freqs_cis = freqs_cis + num_replicated_prefix = 0 + tail_meta = tail_attn_meta( + txt_shard, + encoder_hidden_states.shape[0], + hidden_states.device, + image_seq_len=hidden_states.shape[1], ) - joint_attention_kwargs["attn_mask"] = attn_mask - joint_attention_kwargs["attn_mask_meta"] = attn_mask_meta + if tail_meta is not None: + joint_attention_kwargs = ( + joint_attention_kwargs.copy() if joint_attention_kwargs else {} + ) + joint_attention_kwargs["attn_mask_meta"] = tail_meta + # Single blocks apply RoPE on the relocated [txt_real, img, pad] + # layout, so hand them a cache reordered the same way. + if freqs_cis is not None: + t_loc = txt_shard.local_len + pad = txt_shard.local_pad + singles_freqs_cis = ( + join_seqs(cos[:t_loc], cos[t_loc:], pad, dim=0), + join_seqs(sin[:t_loc], sin[t_loc:], pad, dim=0), + ) if ( joint_attention_kwargs is not None @@ -1229,7 +1163,7 @@ class FluxTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): hidden_states=hidden_states, encoder_hidden_states=encoder_hidden_states, temb=temb, - freqs_cis=freqs_cis, + freqs_cis=singles_freqs_cis, joint_attention_kwargs=joint_attention_kwargs, num_replicated_prefix=num_replicated_prefix, ) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py index 26f8877f7..1409c37ae 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/flux_2.py @@ -23,10 +23,17 @@ from diffusers.models.normalization import AdaLayerNormContinuous from sglang.multimodal_gen.configs.models.dits.flux import FluxConfig from sglang.multimodal_gen.runtime.distributed import ( divide, - get_sp_parallel_rank, - get_sp_world_size, get_tp_world_size, ) +from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import ( + build_shard_plan, + join_seqs, + shard_like, + shard_seq_prefix, + should_shard_text, + split_seqs, + tail_attn_meta, +) from sglang.multimodal_gen.runtime.layers.attention import USPAttention from sglang.multimodal_gen.runtime.layers.layernorm import ( RMSNorm, @@ -60,115 +67,6 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) # pylint: disable=invalid-name -def _shard_text_for_sp( - encoder_hidden_states: torch.Tensor, - freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]], - image_seq_len: int, - num_txt_tokens: int, -) -> Tuple[ - torch.Tensor, - Optional[Tuple[torch.Tensor, torch.Tensor]], - int, - int, - Optional[torch.Tensor], - Optional[Dict[str, int]], -]: - sp_size = get_sp_world_size() - num_replicated_prefix = num_txt_tokens - if sp_size == 1: - return ( - encoder_hidden_states, - freqs_cis, - num_replicated_prefix, - num_txt_tokens, - None, - None, - ) - - sp_rank = get_sp_parallel_rank() - local_txt_tokens = (num_txt_tokens + sp_size - 1) // sp_size - padded_txt_tokens = local_txt_tokens * sp_size - num_pad_tokens = padded_txt_tokens - num_txt_tokens - - if num_pad_tokens > 0: - pad_hidden_states = encoder_hidden_states.new_zeros( - encoder_hidden_states.shape[0], - num_pad_tokens, - encoder_hidden_states.shape[2], - ) - encoder_hidden_states = torch.cat( - [encoder_hidden_states, pad_hidden_states], dim=1 - ) - - encoder_hidden_states = torch.chunk(encoder_hidden_states, sp_size, dim=1)[sp_rank] - if freqs_cis is not None: - cos, sin = freqs_cis - txt_cos = cos[:num_txt_tokens] - txt_sin = sin[:num_txt_tokens] - if num_pad_tokens > 0: - pad_cos = txt_cos.new_ones(num_pad_tokens, txt_cos.shape[1]) - pad_sin = txt_sin.new_zeros(num_pad_tokens, txt_sin.shape[1]) - txt_cos = torch.cat([txt_cos, pad_cos], dim=0) - txt_sin = torch.cat([txt_sin, pad_sin], dim=0) - freqs_cis = ( - torch.cat( - [ - torch.chunk(txt_cos, sp_size, dim=0)[sp_rank], - cos[num_txt_tokens:], - ], - dim=0, - ), - torch.cat( - [ - torch.chunk(txt_sin, sp_size, dim=0)[sp_rank], - sin[num_txt_tokens:], - ], - dim=0, - ), - ) - - num_replicated_prefix = 0 - if num_pad_tokens == 0: - return ( - encoder_hidden_states, - freqs_cis, - num_replicated_prefix, - local_txt_tokens, - None, - None, - ) - - txt_start = sp_rank * local_txt_tokens - valid_txt_tokens = min(local_txt_tokens, max(num_txt_tokens - txt_start, 0)) - text_mask = torch.zeros( - encoder_hidden_states.shape[0], - local_txt_tokens, - dtype=torch.bool, - device=encoder_hidden_states.device, - ) - text_mask[:, :valid_txt_tokens] = True - image_mask = torch.ones( - encoder_hidden_states.shape[0], - image_seq_len, - dtype=torch.bool, - device=encoder_hidden_states.device, - ) - return ( - encoder_hidden_states, - freqs_cis, - num_replicated_prefix, - local_txt_tokens, - torch.cat([text_mask, image_mask], dim=1), - { - "gap_start": (sp_size - 1) * (local_txt_tokens + image_seq_len) - + local_txt_tokens - - num_pad_tokens, - "gap_end": (sp_size - 1) * (local_txt_tokens + image_seq_len) - + local_txt_tokens, - }, - ) - - def _get_qkv_projections( attn: "Flux2Attention", hidden_states, encoder_hidden_states=None ): @@ -465,9 +363,12 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): allow_inplace=True, ) - query = torch.cat([encoder_query, query], dim=1) - key = torch.cat([encoder_key, key], dim=1) - value = torch.cat([encoder_value, value], dim=1) + # join_seqs relocates any SP text tail-pad behind the image (see + # sp_shard.join_seqs for why). + sp_txt_pad = (attn_mask_meta or {}).get("local_pad", 0) + query = join_seqs(encoder_query, query, sp_txt_pad) + key = join_seqs(encoder_key, key, sp_txt_pad) + value = join_seqs(encoder_value, value, sp_txt_pad) else: query, key = apply_qk_norm_with_optional_rope( q=query, @@ -493,12 +394,8 @@ class Flux2Attention(torch.nn.Module, AttentionModuleMixin): hidden_states = hidden_states.to(query.dtype) if encoder_hidden_states is not None: - encoder_hidden_states, hidden_states = hidden_states.split_with_sizes( - [ - encoder_hidden_states.shape[1], - hidden_states.shape[1] - encoder_hidden_states.shape[1], - ], - dim=1, + encoder_hidden_states, hidden_states = split_seqs( + hidden_states, encoder_hidden_states.shape[1], sp_txt_pad ) encoder_hidden_states, _ = self.to_add_out(encoder_hidden_states) @@ -1196,25 +1093,43 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): hidden_states, _ = self.x_embedder(hidden_states) encoder_hidden_states, _ = self.context_embedder(encoder_hidden_states) - ( - encoder_hidden_states, - freqs_cis, - num_replicated_prefix, - num_txt_tokens, - attn_mask, - attn_mask_meta, - ) = _shard_text_for_sp( - encoder_hidden_states, - freqs_cis, - hidden_states.shape[1], - num_txt_tokens, - ) - if attn_mask is not None: - joint_attention_kwargs = ( - joint_attention_kwargs.copy() if joint_attention_kwargs else {} + # Shard the replicated text stream across SP ranks (image latents are + # already sharded); non-divisible lengths tail-pad the last rank and the + # per-request tail meta lets attention skip the pad for free. + num_replicated_prefix = num_txt_tokens + sp_txt_pad = 0 + singles_freqs_cis = freqs_cis + if should_shard_text(num_txt_tokens): + txt_shard = build_shard_plan(num_txt_tokens) + encoder_hidden_states = shard_like(encoder_hidden_states, txt_shard) + if freqs_cis is not None: + cos, sin = freqs_cis + cos = shard_seq_prefix(cos, num_txt_tokens, txt_shard) + sin = shard_seq_prefix(sin, num_txt_tokens, txt_shard) + freqs_cis = (cos, sin) + singles_freqs_cis = freqs_cis + num_replicated_prefix = 0 + num_txt_tokens = txt_shard.local_len + tail_meta = tail_attn_meta( + txt_shard, + encoder_hidden_states.shape[0], + hidden_states.device, + image_seq_len=hidden_states.shape[1], ) - joint_attention_kwargs["attn_mask"] = attn_mask - joint_attention_kwargs["attn_mask_meta"] = attn_mask_meta + if tail_meta is not None: + joint_attention_kwargs = ( + joint_attention_kwargs.copy() if joint_attention_kwargs else {} + ) + joint_attention_kwargs["attn_mask_meta"] = tail_meta + sp_txt_pad = txt_shard.local_pad + # The single-stream trunk applies RoPE on the relocated + # [txt_real, img, pad] layout; reorder its cache to match. + if freqs_cis is not None: + t_loc = txt_shard.local_len + singles_freqs_cis = ( + join_seqs(cos[:t_loc], cos[t_loc:], sp_txt_pad, dim=0), + join_seqs(sin[:t_loc], sin[t_loc:], sp_txt_pad, dim=0), + ) # 4. Double Stream Transformer Blocks for index_block, block in enumerate(self.transformer_blocks): @@ -1227,8 +1142,11 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): joint_attention_kwargs=joint_attention_kwargs, num_replicated_prefix=num_replicated_prefix, ) - # Concatenate text and image streams for single-block inference - hidden_states = torch.cat([encoder_hidden_states, hidden_states], dim=1) + # Concatenate text and image streams for single-block inference; + # join_seqs relocates any SP text tail-pad behind the image once for + # the whole trunk (see sp_shard.join_seqs for why). + txt_real = num_txt_tokens - sp_txt_pad + hidden_states = join_seqs(encoder_hidden_states, hidden_states, sp_txt_pad) # 5. Single Stream Transformer Blocks for index_block, block in enumerate(self.single_transformer_blocks): @@ -1236,13 +1154,14 @@ class Flux2Transformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): hidden_states=hidden_states, encoder_hidden_states=None, temb_mod_params=single_stream_mod, - freqs_cis=freqs_cis, + freqs_cis=singles_freqs_cis, joint_attention_kwargs=joint_attention_kwargs, - text_seq_len=num_txt_tokens, + text_seq_len=txt_real, num_replicated_prefix=num_replicated_prefix, ) - # Remove text tokens from concatenated stream - hidden_states = hidden_states[:, num_txt_tokens:, ...] + # Remove text (and any tail pad) from the concatenated stream + img_end = hidden_states.shape[1] - sp_txt_pad + hidden_states = hidden_states[:, txt_real:img_end, ...] # 6. Output layers hidden_states = self.norm_out(hidden_states, temb) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/mova_video_dit.py b/python/sglang/multimodal_gen/runtime/models/dits/mova_video_dit.py index 5a6a2ded9..fc0f112fb 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/mova_video_dit.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/mova_video_dit.py @@ -153,13 +153,14 @@ class SelfAttention(nn.Module): softmax_scale=None, ) - def forward(self, x, freqs): + def forward(self, x, freqs, attn_mask_meta=None): """ Forward pass for self-attention. Args: x: Input tensor [B, S_local, D] - already sharded by SP when SP > 1 freqs: RoPE frequencies [S_local, 1, head_dim] - should match x's sequence length + attn_mask_meta: sp_shard tail-pad meta; excludes SP padding from attention Returns: Output tensor [B, S_local, D] @@ -189,8 +190,9 @@ class SelfAttention(nn.Module): k = rearrange(k, "b s (n d) -> b s n d", n=self.num_heads_per_rank) v = rearrange(v, "b s (n d) -> b s n d", n=self.num_heads_per_rank) - # USPAttention handles SP communication internally - out = self.attn(q, k, v) + # USPAttention handles SP communication internally; the tail meta keeps + # SP padding out of the softmax. + out = self.attn(q, k, v, attn_mask_meta=attn_mask_meta) out = rearrange(out, "b s n d -> b s (n d)") out, _ = self.o(out) @@ -326,7 +328,7 @@ class DiTBlock(nn.Module): self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) self.mlp_residual = MulAdd() - def forward(self, x, context, t_mod, freqs): + def forward(self, x, context, t_mod, freqs, attn_mask_meta=None): has_seq = len(t_mod.shape) == 4 chunk_dim = 2 if has_seq else 1 # msa: multi-head self-attention mlp: multi-layer perceptron @@ -347,7 +349,9 @@ class DiTBlock(nn.Module): # - layernorm(x) * (1 + scale_msa) + shift_msa input_x = self.norm1(x, shift_msa, scale_msa) # 2. torch.compile may fuse mlp_residual and self_attn_norm - x = self.mlp_residual(self.self_attn(input_x, freqs), gate_msa, x) + x = self.mlp_residual( + self.self_attn(input_x, freqs, attn_mask_meta=attn_mask_meta), gate_msa, x + ) norm_x = self.self_attn_norm(x) # 3. Cross-attention, fuse: # - x = x + 1 * cross_output 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 68ef82384..e65e35de4 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/qwen_image.py @@ -20,10 +20,16 @@ from sglang.multimodal_gen.runtime.distributed import ( get_tp_world_size, ) 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.distributed.sp_shard_utils import ( + build_shard_plan, + join_seqs, + shard_like, + should_shard_text, + split_seqs, + tail_attn_meta, +) from sglang.multimodal_gen.runtime.layers.attention import ( USPAttention, build_varlen_mask_meta, @@ -80,99 +86,6 @@ 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 ): @@ -776,6 +689,8 @@ class QwenImageCrossAttention(nn.Module): # 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) + # Rows of tail padding inside THIS rank's text chunk (sp_shard meta). + sp_txt_pad = (attn_mask_meta or {}).get("local_pad", 0) ( img_query, @@ -834,11 +749,11 @@ class QwenImageCrossAttention(nn.Module): txt_query, txt_key, txt_cache, is_neox=False ) - # Concatenate for joint attention - # Order: [text, image] - joint_query = torch.cat([txt_query, img_query], dim=1) - joint_key = torch.cat([txt_key, img_key], dim=1) - joint_value = torch.cat([txt_value, img_value], dim=1) + # Joint order [text, image]; join_seqs relocates any SP text tail-pad + # behind the image (see sp_shard.join_seqs for why). + joint_query = join_seqs(txt_query, img_query, sp_txt_pad) + joint_key = join_seqs(txt_key, img_key, sp_txt_pad) + joint_value = join_seqs(txt_value, img_value, sp_txt_pad) if attn_mask is None and encoder_hidden_states_mask is not None: image_mask = torch.ones( (hidden_states.shape[0], img_query.shape[1]), @@ -865,8 +780,9 @@ class QwenImageCrossAttention(nn.Module): joint_hidden_states = joint_hidden_states.to(joint_query.dtype) # Split attention outputs back - txt_attn_output = joint_hidden_states[:, :seq_len_txt, :] # Text part - img_attn_output = joint_hidden_states[:, seq_len_txt:, :] # Image part + txt_attn_output, img_attn_output = split_seqs( + joint_hidden_states, seq_len_txt, sp_txt_pad + ) # Apply output projections img_attn_output, _ = self.to_out[0](img_attn_output) @@ -1502,7 +1418,6 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): 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 @@ -1520,29 +1435,24 @@ 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 + elif should_shard_text(encoder_hidden_states.shape[1]): + # Shard the replicated text stream across SP ranks; non-divisible + # lengths tail-pad the last rank and attention skips the pad via the + # per-request tail meta. Otherwise fall through to replicated text. + txt_shard = build_shard_plan(encoder_hidden_states.shape[1]) + encoder_hidden_states = shard_like(encoder_hidden_states, txt_shard) + if freqs_cis is not None: + img_cache, txt_cache = freqs_cis + freqs_cis = (img_cache, shard_like(txt_cache, txt_shard, dim=0)) + tail_meta = tail_attn_meta( + txt_shard, + encoder_hidden_states.shape[0], + hidden_states.device, + image_seq_len=hidden_states.shape[1], ) + if tail_meta is not None: + block_attention_kwargs["attn_mask_meta"] = tail_meta 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) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py index 36de4cf2f..c673fdfcf 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/mova.py @@ -25,13 +25,16 @@ from sglang.multimodal_gen.runtime.distributed import ( ) from sglang.multimodal_gen.runtime.distributed.communication_op import ( cfg_model_parallel_all_reduce, - sequence_model_parallel_all_gather, ) from sglang.multimodal_gen.runtime.distributed.parallel_state import ( get_cfg_group, get_classifier_free_guidance_rank, - get_sp_parallel_rank, - get_sp_world_size, +) +from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import ( + SpShard, + gather_seq, + shard_seq, + tail_attn_meta, ) from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context @@ -651,64 +654,15 @@ class MOVADenoisingStage(PipelineStage): def _shard_sequence_for_sp( self, x: torch.Tensor, dim: int = 1 - ) -> tuple[torch.Tensor, int]: - """ - Shard tensor along sequence dimension for Sequence Parallelism. - - Args: - x: Input tensor - dim: Dimension to shard along - - Returns: - (sharded_tensor, pad_len) - """ - sp_size = get_sp_world_size() - if sp_size <= 1: - return x, 0 - - sp_rank = get_sp_parallel_rank() - seq_len = x.shape[dim] - - # Pad if needed - pad_len = (sp_size - (seq_len % sp_size)) % sp_size - if pad_len > 0: - pad_shape = list(x.shape) - pad_shape[dim] = pad_len - pad = torch.zeros(pad_shape, dtype=x.dtype, device=x.device) - x = torch.cat([x, pad], dim=dim) - - # Shard - chunk_size = x.shape[dim] // sp_size - start = sp_rank * chunk_size - end = start + chunk_size - idx = [slice(None)] * x.dim() - idx[dim] = slice(start, end) - return x[tuple(idx)], pad_len + ) -> tuple[torch.Tensor, SpShard]: + """Tail-padded even shard along the sequence dim (sp_shard.shard_seq).""" + return shard_seq(x, dim=dim) def _gather_sequence_from_sp( - self, x: torch.Tensor, pad_len: int, dim: int = 1 + self, x: torch.Tensor, shard: SpShard, dim: int = 1 ) -> torch.Tensor: - """ - Gather tensor along sequence dimension after Sequence Parallelism. - - Args: - x: Sharded tensor - pad_len: Padding length that was added during sharding - dim: Dimension to gather along - - Returns: - Gathered tensor with padding removed - """ - sp_size = get_sp_world_size() - if sp_size <= 1: - return x - - gathered = sequence_model_parallel_all_gather(x, dim=dim) - if pad_len > 0: - idx = [slice(None)] * gathered.dim() - idx[dim] = slice(0, gathered.shape[dim] - pad_len) - gathered = gathered[tuple(idx)] - return gathered + """Gather an SP-sharded tensor and trim the tail padding.""" + return gather_seq(x, shard.orig_len, dim=dim) def inference_single_step( self, @@ -817,13 +771,20 @@ class MOVADenoisingStage(PipelineStage): ).reshape(full_audio_seq_len, 1, -1) # Shard sequences for SP - visual_x, visual_pad_len = self._shard_sequence_for_sp(visual_x, dim=1) - audio_x, audio_pad_len = self._shard_sequence_for_sp(audio_x, dim=1) + visual_x, visual_shard = self._shard_sequence_for_sp(visual_x, dim=1) + audio_x, audio_shard = self._shard_sequence_for_sp(audio_x, dim=1) # Shard freqs to match local sequence length visual_freqs, _ = self._shard_sequence_for_sp(visual_freqs, dim=0) audio_freqs, _ = self._shard_sequence_for_sp(audio_freqs, dim=0) + # Tail-pad meta so self-attention excludes SP padding (built once per + # step, shared by every block). + visual_attn_meta = tail_attn_meta( + visual_shard, visual_x.shape[0], visual_x.device + ) + audio_attn_meta = tail_attn_meta(audio_shard, audio_x.shape[0], audio_x.device) + # Forward through dual-tower DiT visual_x, audio_x = self.forward_dual_tower_dit( visual_dit=visual_dit, @@ -839,11 +800,13 @@ class MOVADenoisingStage(PipelineStage): video_fps=video_fps, full_visual_seq_len=full_visual_seq_len, full_audio_seq_len=full_audio_seq_len, + visual_attn_meta=visual_attn_meta, + audio_attn_meta=audio_attn_meta, ) # Gather sequences back from SP before head/unpatchify - visual_x = self._gather_sequence_from_sp(visual_x, visual_pad_len, dim=1) - audio_x = self._gather_sequence_from_sp(audio_x, audio_pad_len, dim=1) + visual_x = self._gather_sequence_from_sp(visual_x, visual_shard, dim=1) + audio_x = self._gather_sequence_from_sp(audio_x, audio_shard, dim=1) visual_output = visual_dit.head(visual_x, visual_t) visual_output = visual_dit.unpatchify(visual_output, grid_size) @@ -871,6 +834,8 @@ class MOVADenoisingStage(PipelineStage): condition_scale: float | None = 1.0, a2v_condition_scale: float | None = None, v2a_condition_scale: float | None = None, + visual_attn_meta: dict | None = None, + audio_attn_meta: dict | None = None, ): """ Forward pass through dual-tower DiT with cross-modal interaction. @@ -933,15 +898,29 @@ class MOVADenoisingStage(PipelineStage): # Self-attention and FFN in DiT blocks visual_x = visual_block( - visual_x, visual_context, visual_t_mod, visual_freqs + visual_x, + visual_context, + visual_t_mod, + visual_freqs, + attn_mask_meta=visual_attn_meta, + ) + audio_x = audio_block( + audio_x, + audio_context, + audio_t_mod, + audio_freqs, + attn_mask_meta=audio_attn_meta, ) - audio_x = audio_block(audio_x, audio_context, audio_t_mod, audio_freqs) # Process remaining visual layers (if visual has more layers than audio) for layer_idx in range(min_layers, visual_layers): visual_block = visual_dit.blocks[layer_idx] visual_x = visual_block( - visual_x, visual_context, visual_t_mod, visual_freqs + visual_x, + visual_context, + visual_t_mod, + visual_freqs, + attn_mask_meta=visual_attn_meta, ) return visual_x, audio_x diff --git a/python/sglang/multimodal_gen/test/server/consistency_thresholds/h100.json b/python/sglang/multimodal_gen/test/server/consistency_thresholds/h100.json index 29705f212..62a5d8554 100644 --- a/python/sglang/multimodal_gen/test/server/consistency_thresholds/h100.json +++ b/python/sglang/multimodal_gen/test/server/consistency_thresholds/h100.json @@ -229,12 +229,6 @@ "psnr_threshold": 30.0, "mean_abs_diff_threshold": 5.6 }, - "mova_360p_ring1_uly2": { - "clip_threshold": 0.96, - "ssim_threshold": 0.91, - "psnr_threshold": 30.0, - "mean_abs_diff_threshold": 6.2 - }, "wan2_1_i2v_14b_lora_2gpu": { "clip_threshold": 0.97, "ssim_threshold": 0.90, diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 3fde45b97..d72b8eb4d 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -34,7 +34,7 @@ if TYPE_CHECKING: logger = init_logger(__name__) -SGL_TEST_FILES_CI_DATA_REVISION = "c28c11c24125b009f9e602fd5ae0d0ddcdc11d36" +SGL_TEST_FILES_CI_DATA_REVISION = "46b9b53a429606cb6739c861f275c1277c314a10" if current_platform.is_npu(): SGL_TEST_FILES_CI_DATA_REVISION = "670d66a8a290b62c0c3c077b3e9b0f4a4d9a44e7" diff --git a/python/sglang/multimodal_gen/test/unit/test_sp_shard.py b/python/sglang/multimodal_gen/test/unit/test_sp_shard.py new file mode 100644 index 000000000..f16173426 --- /dev/null +++ b/python/sglang/multimodal_gen/test/unit/test_sp_shard.py @@ -0,0 +1,213 @@ +"""Unit tests for the unified SP shard helpers (pure logic, no distributed).""" + +import pytest +import torch + +from sglang.multimodal_gen.runtime.distributed import sp_shard_utils as sps +from sglang.multimodal_gen.runtime.distributed.sp_shard_utils import ( + SpShard, + shard_like, + tail_attn_meta, +) + + +def _fake_sp(monkeypatch, sp_size, sp_rank=0, ring=1): + monkeypatch.setattr(sps, "get_sp_world_size", lambda: sp_size) + monkeypatch.setattr(sps, "get_sp_parallel_rank", lambda: sp_rank) + monkeypatch.setattr(sps, "get_ring_parallel_world_size", lambda: ring) + + +# --- build_shard_plan math -------------------------------------------------------- + + +def test_plan_shard_divisible(monkeypatch): + _fake_sp(monkeypatch, 2, 1) + s = sps.build_shard_plan(16) + assert (s.local_len, s.num_pad, s.local_pad) == (8, 0, 0) + + +def test_plan_shard_padded_last_rank(monkeypatch): + _fake_sp(monkeypatch, 4, 3) + s = sps.build_shard_plan(14) + assert (s.local_len, s.num_pad) == (4, 2) + assert s.local_pad == 2 and s.local_real_len == 2 + + +def test_plan_shard_pad_only_on_last_rank(monkeypatch): + _fake_sp(monkeypatch, 4, 0) + s = sps.build_shard_plan(14) + assert s.local_pad == 0 and s.local_real_len == 4 + + +def test_plan_shard_sp1_noop(monkeypatch): + _fake_sp(monkeypatch, 1) + s = sps.build_shard_plan(15) + assert (s.local_len, s.num_pad, s.sp_size) == (15, 0, 1) + + +# --- shard_like ------------------------------------------------------------- + + +def test_shard_like_zero_pads_tail(): + shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1) + x = torch.arange(15, dtype=torch.float32).unsqueeze(0).unsqueeze(-1) + local = shard_like(x, shard, dim=1) + assert local.shape[1] == 8 + assert local[0, -1, 0].item() == 0.0 # tail pad + assert local[0, 0, 0].item() == 8.0 # rank1 starts at token 8 + + +def test_shard_like_repeat_last(): + shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1) + x = torch.arange(15, dtype=torch.float32).unsqueeze(-1) + local = shard_like(x, shard, dim=0, pad_mode="repeat_last") + assert local[-1, 0].item() == 14.0 # repeated last row, not zero + + +def test_shard_like_chunks_align_across_tensors(): + # RoPE cache sharded with the same plan stays aligned with hidden states. + shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=0) + x = torch.arange(15).unsqueeze(0).unsqueeze(-1).float() + rope = torch.arange(15).unsqueeze(-1).float() + assert torch.equal( + shard_like(x, shard, dim=1)[0, :, 0], shard_like(rope, shard, dim=0)[:, 0] + ) + + +# --- tail_attn_meta --------------------------------------------------------- + + +def test_tail_meta_none_when_divisible(): + shard = SpShard(orig_len=16, local_len=8, num_pad=0, sp_size=2, sp_rank=0) + assert tail_attn_meta(shard, 1, torch.device("cpu")) is None + + +def test_tail_meta_single_stream(): + shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1) + meta = tail_attn_meta(shard, 1, torch.device("cpu")) + assert meta["pad_start"] == 15 and meta["pad_end"] == 16 + assert meta["local_pad"] == 1 + assert meta["cu_seqlens_tail"].tolist() == [0, 15, 16] + assert meta["max_seqlen_tail"] == 15 + + +def test_tail_meta_joint_layout_and_batch(): + # sp=2, local_txt=8 (1 pad), img=100 per rank -> S = 2*(8+100) = 216. + shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1) + meta = tail_attn_meta(shard, 2, torch.device("cpu"), image_seq_len=100) + assert meta["pad_start"] == 215 and meta["pad_end"] == 216 + assert meta["cu_seqlens_tail"].tolist() == [0, 215, 216, 431, 432] + + +def test_tail_meta_max_seqlen_covers_pad_segment(): + # Degenerate short sequence: num_pad (3) > valid (1). FA requires + # max_seqlen >= the longest segment, i.e. the pad block here. + shard = SpShard(orig_len=1, local_len=1, num_pad=3, sp_size=4, sp_rank=3) + meta = tail_attn_meta(shard, 1, torch.device("cpu")) + assert meta["max_seqlen_tail"] == 3 + + +def test_tail_meta_matches_legacy_gap_formula(): + # The tail layout puts the pad exactly where the legacy per-model gap + # formula pointed, minus the relocation: end == S (global tail). + sp, local_txt, img, num_pad = 3, 5, 40, 2 + shard = SpShard( + orig_len=sp * local_txt - num_pad, + local_len=local_txt, + num_pad=num_pad, + sp_size=sp, + sp_rank=sp - 1, + ) + meta = tail_attn_meta(shard, 1, torch.device("cpu"), image_seq_len=img) + seq = sp * (local_txt + img) + assert meta["pad_end"] == seq + assert meta["pad_start"] == seq - num_pad + + +# --- plan_text_strategy ----------------------------------------------------- + + +def test_strategy_sp1_replicates(monkeypatch): + _fake_sp(monkeypatch, 1) + assert sps.plan_text_strategy(100) == "replicate" + + +def test_strategy_shard_when_legal(monkeypatch): + _fake_sp(monkeypatch, 2) + assert sps.plan_text_strategy(15) == "shard" + assert sps.plan_text_strategy(16) == "shard" + + +def test_strategy_ring_blocks_padded_shard(monkeypatch): + _fake_sp(monkeypatch, 2, ring=2) + assert sps.plan_text_strategy(15) == "replicate" # padded shard needs mask + assert sps.plan_text_strategy(16) == "shard" # divisible: no mask needed + + +def test_strategy_min_len_threshold(monkeypatch): + _fake_sp(monkeypatch, 2) + monkeypatch.setattr(sps, "_TEXT_SHARD_MIN", 64) + assert sps.plan_text_strategy(32) == "replicate" + assert sps.plan_text_strategy(64) == "shard" + + +# --- join_seqs / split_seqs / shard_seq_prefix ------------------------------ + + +def test_join_split_roundtrip_with_pad(): + # Joint [text, image] with 2 tail-pad rows relocated behind the image. + txt = torch.arange(6, dtype=torch.float32).view(1, 6, 1) # rows 4,5 are pad + img = (torch.arange(3, dtype=torch.float32) + 100).view(1, 3, 1) + joint = sps.join_seqs(txt, img, local_pad=2) + assert joint[0, :, 0].tolist() == [0, 1, 2, 3, 100, 101, 102, 4, 5] + txt_back, img_back = sps.split_seqs(joint, prefix_len=6, local_pad=2) + assert torch.equal(txt_back, txt) and torch.equal(img_back, img) + + +def test_join_split_roundtrip_no_pad(): + txt = torch.randn(1, 4, 2) + img = torch.randn(1, 3, 2) + joint = sps.join_seqs(txt, img, local_pad=0) + assert torch.equal(joint, torch.cat([txt, img], dim=1)) + txt_back, img_back = sps.split_seqs(joint, prefix_len=4, local_pad=0) + assert torch.equal(txt_back, txt) and torch.equal(img_back, img) + + +def test_shard_seq_prefix_only_touches_prefix(): + # Joint RoPE cache [txt(15); img(4)]: text segment shards, image stays. + shard = SpShard(orig_len=15, local_len=8, num_pad=1, sp_size=2, sp_rank=1) + cache = torch.arange(19, dtype=torch.float32).unsqueeze(-1) + out = sps.shard_seq_prefix(cache, 15, shard, dim=0) + assert out.shape[0] == 8 + 4 + assert out[0, 0].item() == 8.0 # rank1 text chunk starts at token 8 + assert out[-4:, 0].flatten().tolist() == [15, 16, 17, 18] # image untouched + + +def test_should_shard_text_gate(monkeypatch): + _fake_sp(monkeypatch, 2) + assert sps.should_shard_text(15) is True + _fake_sp(monkeypatch, 1) + assert sps.should_shard_text(15) is False + + +# --- gather_seq ------------------------------------------------------------- + + +def test_gather_seq_sp1_noop(monkeypatch): + _fake_sp(monkeypatch, 1) + x = torch.randn(1, 5, 2) + assert sps.gather_seq(x, 5, dim=1) is x + + +def test_gather_seq_trims(monkeypatch): + _fake_sp(monkeypatch, 2) + monkeypatch.setattr( + sps, "sequence_model_parallel_all_gather", lambda t, dim: torch.cat([t, t], dim) + ) + local = torch.randn(1, 8, 2) + out = sps.gather_seq(local, 15, dim=1) + assert out.shape[1] == 15 + + +if __name__ == "__main__": + pytest.main([__file__, "-q"])