[diffusion] optimize: reduce cosmos3 denoise overhead (#26973)

This commit is contained in:
Mick
2026-06-02 14:23:02 +08:00
committed by GitHub
parent 3b26644bc4
commit 1033d835ff
3 changed files with 98 additions and 45 deletions
@@ -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)
@@ -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.
@@ -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