From 1033d835ff293fdb969073f1291da250c8070745 Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 2 Jun 2026 14:23:02 +0800 Subject: [PATCH] [diffusion] optimize: reduce cosmos3 denoise overhead (#26973) --- .../runtime/layers/attention/layer.py | 44 +++++++++- .../runtime/models/dits/cosmos3video.py | 82 ++++++++++--------- .../stages/model_specific_stages/cosmos3.py | 17 +++- 3 files changed, 98 insertions(+), 45 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py index cc4515756..b32cae472 100644 --- a/python/sglang/multimodal_gen/runtime/layers/attention/layer.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/layer.py @@ -750,6 +750,32 @@ class USPAttention(nn.Module): return torch.cat([out_rep, out_shard], dim=1) + def forward_with_replicated_kv_prefix( + self, + q: torch.Tensor, + k_prefix: torch.Tensor, + v_prefix: torch.Tensor, + k_suffix: torch.Tensor, + v_suffix: torch.Tensor, + ) -> torch.Tensor: + """attention with replicated K/V prefix supplied separately""" + forward_context: ForwardContext = get_forward_context() + ctx_attn_metadata = forward_context.attn_metadata + + if self.skip_sequence_parallel or get_sequence_parallel_world_size() == 1: + k = torch.cat([k_prefix, k_suffix], dim=1) + v = torch.cat([v_prefix, v_suffix], dim=1) + return self.attn_impl.forward(q, k, v, ctx_attn_metadata) + + if get_ulysses_parallel_world_size() == 1: + k = torch.cat([k_prefix, k_suffix], dim=1) + v = torch.cat([v_prefix, v_suffix], dim=1) + return self(q, k, v) + + return self._forward_with_replicated_kv_prefix_split( + q, k_prefix, v_prefix, k_suffix, v_suffix, ctx_attn_metadata + ) + def _forward_with_replicated_kv_prefix( self, q: torch.Tensor, @@ -771,11 +797,25 @@ class USPAttention(nn.Module): 3. Concatenate prefix + suffix on the sequence dim and attend. 4. All-to-all the output back (head shard → seq shard). """ - sp_rank = get_sp_parallel_rank() - k_rep, k_shard = k[:, :num_rep], k[:, num_rep:] v_rep, v_shard = v[:, :num_rep], v[:, num_rep:] + return self._forward_with_replicated_kv_prefix_split( + q, k_rep, v_rep, k_shard, v_shard, ctx_attn_metadata + ) + + def _forward_with_replicated_kv_prefix_split( + self, + q: torch.Tensor, + k_rep: torch.Tensor, + v_rep: torch.Tensor, + k_shard: torch.Tensor, + v_shard: torch.Tensor, + ctx_attn_metadata, + ) -> torch.Tensor: + """split form avoids materializing full K/V before Ulysses all-to-all""" + sp_rank = get_sp_parallel_rank() + q = _usp_input_all_to_all(q, head_dim=2) k_shard = _usp_input_all_to_all(k_shard, head_dim=2) v_shard = _usp_input_all_to_all(v_shard, head_dim=2) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py b/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py index 46bdd8318..69d0e9926 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/cosmos3video.py @@ -428,18 +428,21 @@ class Cosmos3CausalAttention(nn.Module): batch_size, seq_len = hidden_states.shape[:2] qkv, _ = self.to_qkv(hidden_states) - # split returns strided views into qkv; .contiguous() before .view() - # because the per-head reshape needs row-major memory. - q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - q = q.contiguous().view( - batch_size, seq_len, self.num_attention_heads, self.head_dim - ) - k = k.contiguous().view( - batch_size, seq_len, self.num_key_value_heads, self.head_dim - ) - v = v.contiguous().view( - batch_size, seq_len, self.num_key_value_heads, self.head_dim + qkv = qkv.view( + batch_size, + seq_len, + self.num_attention_heads + 2 * self.num_key_value_heads, + self.head_dim, ) + q = qkv[:, :, : self.num_attention_heads, :] + k = qkv[ + :, + :, + self.num_attention_heads : self.num_attention_heads + + self.num_key_value_heads, + :, + ] + v = qkv[:, :, self.num_attention_heads + self.num_key_value_heads :, :] q = F.rms_norm( q, (self.head_dim,), self.norm_q.weight, self.norm_q.variance_epsilon @@ -536,16 +539,21 @@ class Cosmos3CrossAttention(nn.Module): batch_size, seq_len_gen = hidden_states.shape[:2] qkv, _ = self.to_qkv(hidden_states) - q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) - q = q.contiguous().view( - batch_size, seq_len_gen, self.num_attention_heads, self.head_dim - ) - k = k.contiguous().view( - batch_size, seq_len_gen, self.num_key_value_heads, self.head_dim - ) - v = v.contiguous().view( - batch_size, seq_len_gen, self.num_key_value_heads, self.head_dim + qkv = qkv.view( + batch_size, + seq_len_gen, + self.num_attention_heads + 2 * self.num_key_value_heads, + self.head_dim, ) + q = qkv[:, :, : self.num_attention_heads, :] + k = qkv[ + :, + :, + self.num_attention_heads : self.num_attention_heads + + self.num_key_value_heads, + :, + ] + v = qkv[:, :, self.num_attention_heads + self.num_key_value_heads :, :] q = F.rms_norm( q, (self.head_dim,), self.norm_q.weight, self.norm_q.variance_epsilon @@ -558,10 +566,7 @@ class Cosmos3CrossAttention(nn.Module): # K/V = [text (replicated full on every SP rank) | image (sharded same as Q)]. # USPAttention routes through the registered attention backend (FA, sage, # …) and handles the Ulysses all-to-all when SP > 1. - num_und = k_und.shape[1] - k = torch.cat([k_und, k], dim=1) - v = torch.cat([v_und, v], dim=1) - out = self.attn(q, k, v, num_replicated_kv_prefix=num_und) + out = self.attn.forward_with_replicated_kv_prefix(q, k_und, v_und, k, v) out = out.reshape(batch_size, seq_len_gen, -1) out, _ = self.to_out(out) return out @@ -1158,25 +1163,24 @@ class Cosmos3OmniTransformer(CachableDiT): self.cached_kv[cache_key] = self.language_model( text_ids, text_mask, freqs_und[0], freqs_und[1] ) - self.cached_freqs_gen[cache_key] = freqs_gen + cos_gen, sin_gen = freqs_gen + if sequence_shard_enabled: + if seq_shard_pad > 0: + pad_cos = cos_gen[:, -1:].expand(-1, seq_shard_pad, -1) + pad_sin = sin_gen[:, -1:].expand(-1, seq_shard_pad, -1) + cos_gen = torch.cat([cos_gen, pad_cos], dim=1) + sin_gen = torch.cat([sin_gen, pad_sin], dim=1) + cos_gen = cos_gen.view(batch_size, self.sp_size, local_seq_len, -1) + sin_gen = sin_gen.view(batch_size, self.sp_size, local_seq_len, -1) + cos_gen = cos_gen[:, self.sp_rank, :, :] + sin_gen = sin_gen[:, self.sp_rank, :, :] + cos_gen = cos_gen.unsqueeze(2) # [B, S, 1, D] + sin_gen = sin_gen.unsqueeze(2) + self.cached_freqs_gen[cache_key] = (cos_gen, sin_gen) freqs_gen = self.cached_freqs_gen[cache_key] cos_gen, sin_gen = freqs_gen - if sequence_shard_enabled: - if seq_shard_pad > 0: - pad_cos = cos_gen[:, -1:].expand(-1, seq_shard_pad, -1) - pad_sin = sin_gen[:, -1:].expand(-1, seq_shard_pad, -1) - cos_gen = torch.cat([cos_gen, pad_cos], dim=1) - sin_gen = torch.cat([sin_gen, pad_sin], dim=1) - cos_gen = cos_gen.view(batch_size, self.sp_size, local_seq_len, -1) - sin_gen = sin_gen.view(batch_size, self.sp_size, local_seq_len, -1) - cos_gen = cos_gen[:, self.sp_rank, :, :] - sin_gen = sin_gen[:, self.sp_rank, :, :] - - cos_gen = cos_gen.unsqueeze(2) # [B, S, 1, D] - sin_gen = sin_gen.unsqueeze(2) - # Run GEN layers. `residual` is threaded so each layer's # input_layernorm and post_attention_layernorm can use the # fused add+rmsnorm path instead of separate add + norm kernels. diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/cosmos3.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/cosmos3.py index 54f1c03f1..9afd4d5a5 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/cosmos3.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/cosmos3.py @@ -499,6 +499,7 @@ class Cosmos3DenoisingStage(PipelineStage): cache_key: str = "default", noisy_frame_mask: torch.Tensor | None = None, max_text_seq_len: int | None = None, + current_timestep: int | None = None, ) -> torch.Tensor: """Run transformer forward pass. @@ -513,10 +514,9 @@ class Cosmos3DenoisingStage(PipelineStage): and "uncond" for unconditional to enable cache reuse across steps. noisy_frame_mask: Optional [B, 1, T, 1, 1] I2V conditioning mask. """ - with set_forward_context( - current_timestep=int(timestep.flatten()[0].item()), - attn_metadata=None, - ): + if current_timestep is None: + current_timestep = int(timestep.flatten()[0].item()) + with set_forward_context(current_timestep=current_timestep, attn_metadata=None): return self.transformer( hidden_states=latents, encoder_hidden_states=None, # Not used by Cosmos3 @@ -641,6 +641,7 @@ class Cosmos3DenoisingStage(PipelineStage): noisy_frame_mask=velocity_mask, cond_text_seq_len=batch.extra["cond_text_seq_len"], uncond_text_seq_len=batch.extra["uncond_text_seq_len"], + current_timestep=i, ) elif effective_scale == 1.0: noise_pred = self._run_transformer( @@ -653,6 +654,7 @@ class Cosmos3DenoisingStage(PipelineStage): cache_key="cond", noisy_frame_mask=velocity_mask, max_text_seq_len=batch.extra["cond_text_seq_len"], + current_timestep=i, ) else: noise_pred = self._predict_noise_cfg_batched( @@ -670,6 +672,7 @@ class Cosmos3DenoisingStage(PipelineStage): batch.extra["cond_text_seq_len"], batch.extra["uncond_text_seq_len"], ), + current_timestep=i, ) else: noise_pred = self._run_transformer( @@ -682,6 +685,7 @@ class Cosmos3DenoisingStage(PipelineStage): cache_key="cond", noisy_frame_mask=velocity_mask, max_text_seq_len=batch.extra["cond_text_seq_len"], + current_timestep=i, ) # I2V: zero-velocity at conditioned frames so the scheduler keeps @@ -717,6 +721,7 @@ class Cosmos3DenoisingStage(PipelineStage): guidance_scale: float, noisy_frame_mask: torch.Tensor | None = None, max_text_seq_len: int | None = None, + current_timestep: int | None = None, ) -> torch.Tensor: """Run CFG by stacking both branches into a batch_size=2 forward. @@ -744,6 +749,7 @@ class Cosmos3DenoisingStage(PipelineStage): cache_key="cfg_batched", noisy_frame_mask=mask_batched, max_text_seq_len=max_text_seq_len, + current_timestep=current_timestep, ) noise_pred_uncond, noise_pred_cond = noise_pred.chunk(2, dim=0) @@ -767,6 +773,7 @@ class Cosmos3DenoisingStage(PipelineStage): noisy_frame_mask: torch.Tensor | None = None, cond_text_seq_len: int | None = None, uncond_text_seq_len: int | None = None, + current_timestep: int | None = None, ) -> torch.Tensor: """Run CFG with one branch per CFG rank, combined by all-reduce. @@ -787,6 +794,7 @@ class Cosmos3DenoisingStage(PipelineStage): cache_key="cond", noisy_frame_mask=noisy_frame_mask, max_text_seq_len=cond_text_seq_len, + current_timestep=current_timestep, ) partial = guidance_scale * noise_pred else: @@ -800,6 +808,7 @@ class Cosmos3DenoisingStage(PipelineStage): cache_key="uncond", noisy_frame_mask=noisy_frame_mask, max_text_seq_len=uncond_text_seq_len, + current_timestep=current_timestep, ) partial = (1.0 - guidance_scale) * noise_pred