[diffusion] fix: skip USP for cross-attention with replicated KV for wan (#19419)

Co-authored-by: Mick <mickjagger19@icloud.com>
This commit is contained in:
AichenF
2026-03-02 19:52:08 +08:00
committed by GitHub
parent da2a0240f7
commit 8df9b8dce9
2 changed files with 20 additions and 20 deletions

View File

@@ -22,7 +22,6 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.attention_backend i
)
from sglang.multimodal_gen.runtime.layers.attention.selector import get_attn_backend
from sglang.multimodal_gen.runtime.layers.usp import (
_ulysses_input_split,
_usp_input_all_to_all,
_usp_output_all_to_all,
ring_attn,
@@ -304,9 +303,17 @@ class USPAttention(nn.Module):
supported_attention_backends: set[AttentionBackendEnum] | None = None,
prefix: str = "",
dropout_rate: float = 0.0,
is_cross_attention: bool = False,
skip_sequence_parallel: bool = False,
**extra_impl_args,
) -> None:
"""
Args:
skip_sequence_parallel:
when KV is replicated across all SP ranks (e.g. cross-attention to
text/image encoder outputs), the full USP pipeline is redundant:
each rank's local Q shard can attend directly to the locally-held
full KV without any collective communication.
"""
super().__init__()
if softmax_scale is None:
self.softmax_scale = head_size**-0.5
@@ -337,7 +344,8 @@ class USPAttention(nn.Module):
self.dtype = dtype
self.causal = causal
self.dropout_p = dropout_rate
self.is_cross_attention = is_cross_attention
self.skip_sequence_parallel = skip_sequence_parallel
def forward(
self,
@@ -354,13 +362,16 @@ class USPAttention(nn.Module):
q, k, v: [B, S_local, H, D]
Note: Replicated tensors are not supported in this implementation.
When skip_sequence_parallel=True (set at construction time), all SP
communication is bypassed — use this for cross-attention where KV
content is replicated across ranks (distinct from replicated_k/v args).
"""
assert (
replicated_q is None and replicated_k is None and replicated_v is None
), "USPAttention does not support replicated_qkv."
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
if get_sequence_parallel_world_size() == 1:
if self.skip_sequence_parallel or get_sequence_parallel_world_size() == 1:
# No sequence parallelism, just run local attention.
out = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
return out
@@ -369,15 +380,11 @@ class USPAttention(nn.Module):
if get_ulysses_parallel_world_size() > 1:
# -> [B, S, H_local, D]
q = _usp_input_all_to_all(q, head_dim=2)
if self.is_cross_attention:
k = _ulysses_input_split(k, dim=2)
v = _ulysses_input_split(v, dim=2)
else:
k = _usp_input_all_to_all(k, head_dim=2)
v = _usp_input_all_to_all(v, head_dim=2)
k = _usp_input_all_to_all(k, head_dim=2)
v = _usp_input_all_to_all(v, head_dim=2)
# Ring Attention within subgroups or local attention
if get_ring_parallel_world_size() > 1 and not self.is_cross_attention:
if get_ring_parallel_world_size() > 1:
out = ring_attn(
q,
k,

View File

@@ -174,27 +174,20 @@ class WanSelfAttention(nn.Module):
softmax_scale=None,
causal=False,
supported_attention_backends=supported_attention_backends,
is_cross_attention=is_cross_attention,
skip_sequence_parallel=is_cross_attention,
)
def forward(self, x: torch.Tensor, context: torch.Tensor, context_lens: int):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
seq_lens(Tensor): Shape [B]
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
pass
class WanT2VCrossAttention(WanSelfAttention):
def __init__(self, *args, **kwargs):
super().__init__(
*args,
**kwargs,
is_cross_attention=True,
)
super().__init__(*args, **kwargs, is_cross_attention=True)
def forward(self, x, context, context_lens):
r"""