This commit is contained in:
kijai
2025-07-12 02:06:02 +03:00
parent 127bfa2c99
commit 039b211ef7
2 changed files with 42 additions and 82 deletions
+2 -26
View File
@@ -181,30 +181,6 @@ def attention(
version=fa_version,
)
elif attention_mode == 'sdpa':
# if q_lens is not None or k_lens is not None:
# warnings.warn(
# 'Padding mask is disabled when using scaled_dot_product_attention. It can have a significant impact on performance.'
# )
attn_mask = None
q = q.transpose(1, 2)#.to(dtype)
k = k.transpose(1, 2)#.to(dtype)
v = v.transpose(1, 2)#.to(dtype)
out = torch.nn.functional.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
out = out.transpose(1, 2).contiguous()
return out
return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous()
elif attention_mode == 'sageattn':
attn_mask = None
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
out = sageattn_func(
q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p)
out = out.transpose(1, 2).contiguous()
return out
return sageattn_func(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous()
+40 -56
View File
@@ -212,7 +212,6 @@ class WanSelfAttention(nn.Module):
in_features,
out_features,
num_heads,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
attention_mode='sdpa'):
@@ -221,7 +220,6 @@ class WanSelfAttention(nn.Module):
self.dim = out_features
self.num_heads = num_heads
self.head_dim = out_features // num_heads
self.window_size = window_size
self.qk_norm = qk_norm
self.eps = eps
self.attention_mode = attention_mode
@@ -241,7 +239,7 @@ class WanSelfAttention(nn.Module):
v = self.v(x).view(b, s, n, d)
return q, k, v
def forward(self, q, k, v, seq_lens, grid_sizes, freqs, block_mask=None):
def forward(self, q, k, v, seq_lens, block_mask=None):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -277,7 +275,6 @@ class WanSelfAttention(nn.Module):
x = attention(
q, k, v,
k_lens=seq_lens,
window_size=self.window_size,
attention_mode=self.attention_mode
)
@@ -336,7 +333,6 @@ class WanSelfAttention(nn.Module):
k=chunk_k,
v=chunk_v,
k_lens=seq_lens,
window_size=self.window_size,
attention_mode=self.attention_mode)
outputs.append(chunk_out)
@@ -350,7 +346,6 @@ class WanSelfAttention(nn.Module):
k=k,
v=v,
k_lens=seq_lens,
window_size=self.window_size,
attention_mode=self.attention_mode)
# output
@@ -396,8 +391,8 @@ class WanSelfAttention(nn.Module):
#region T2V crossattn
class WanT2VCrossAttention(WanSelfAttention):
def __init__(self, in_features, out_features, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(in_features, out_features, num_heads, window_size, qk_norm, eps)
def __init__(self, in_features, out_features, num_heads, qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(in_features, out_features, num_heads, qk_norm, eps)
self.attention_mode = attention_mode
def forward(self, x, context, context_lens, clip_embed=None, audio_proj=None, audio_context_lens=None, audio_scale=1.0,
@@ -439,8 +434,8 @@ class WanT2VCrossAttention(WanSelfAttention):
class WanI2VCrossAttention(WanSelfAttention):
def __init__(self, in_features, out_features, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(in_features, out_features, num_heads, window_size, qk_norm, eps)
def __init__(self, in_features, out_features, num_heads, qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(in_features, out_features, num_heads, qk_norm, eps)
self.k_img = nn.Linear(in_features, out_features)
self.v_img = nn.Linear(in_features, out_features)
self.norm_k_img = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
@@ -511,7 +506,6 @@ class WanAttentionBlock(nn.Module):
ffn_dim,
ffn2_dim,
num_heads,
window_size=(-1, -1),
qk_norm=True,
cross_attn_norm=False,
eps=1e-6,
@@ -522,7 +516,6 @@ class WanAttentionBlock(nn.Module):
self.dim = out_features
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.window_size = window_size
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
@@ -531,7 +524,7 @@ class WanAttentionBlock(nn.Module):
# layers
self.norm1 = WanLayerNorm(out_features, eps)
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, window_size, qk_norm,
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm,
eps, self.attention_mode)
if cross_attn_type != "no_cross_attn":
self.norm3 = WanLayerNorm(
@@ -540,7 +533,6 @@ class WanAttentionBlock(nn.Module):
self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](in_features,
out_features,
num_heads,
(-1, -1),
qk_norm,
eps,#attention_mode=attention_mode sageattn doesn't seem faster here
)
@@ -552,7 +544,7 @@ class WanAttentionBlock(nn.Module):
# modulation
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
#@torch.compiler.disable()
@torch.compiler.disable()
def get_mod(self, e):
if e.dim() == 3:
modulation = self.modulation # 1, 6, dim
@@ -635,7 +627,8 @@ class WanAttentionBlock(nn.Module):
k=rope_apply(k, grid_sizes, freqs)
#self-attention
if context is not None and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1:
split_attn = context is not None and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1
if split_attn:
y = self.self_attn.forward_split(
q, k, v,
seq_lens, grid_sizes, freqs,
@@ -644,11 +637,7 @@ class WanAttentionBlock(nn.Module):
video_attention_split_steps=video_attention_split_steps
)
else:
y = self.self_attn.forward(
q, k, v,
seq_lens, grid_sizes,
freqs, block_mask=block_mask
)
y = self.self_attn.forward(q, k, v, seq_lens, block_mask=block_mask)
#multitalk mask
if ref_target_masks is not None:
@@ -667,7 +656,7 @@ class WanAttentionBlock(nn.Module):
# cross-attention & ffn function
if context is not None:
if (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1:
if split_attn:
if nag_context is not None:
raise NotImplementedError("nag_context is not supported in split_cross_attn_ffn")
x = self.split_cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes)
@@ -765,7 +754,6 @@ class VaceWanAttentionBlock(WanAttentionBlock):
ffn_dim,
ffn2_dim,
num_heads,
window_size=(-1, -1),
qk_norm=True,
cross_attn_norm=False,
eps=1e-6,
@@ -773,25 +761,14 @@ class VaceWanAttentionBlock(WanAttentionBlock):
attention_mode='sdpa',
rope_func="comfy"
):
super().__init__(cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, window_size, qk_norm, cross_attn_norm, eps, attention_mode, rope_func)
super().__init__(cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm, cross_attn_norm, eps, attention_mode, rope_func)
self.block_id = block_id
if block_id == 0:
self.before_proj = nn.Linear(in_features, out_features)
self.after_proj = nn.Linear(in_features, out_features)
def forward(self, c_list, x, intermediate_device=None, nonblocking=True, **kwargs):
if self.block_id == 0:
c = self.before_proj(c_list[0]) + x
all_c = []
else:
all_c = c_list
c = all_c.pop(-1)
c = super().forward(c, **kwargs)
c_skip = self.after_proj(c)
all_c += [c_skip.to(intermediate_device, non_blocking=nonblocking), c]
return all_c
def forward(self, c, **kwargs):
return super().forward(c, **kwargs)
class BaseWanAttentionBlock(WanAttentionBlock):
def __init__(
@@ -802,7 +779,6 @@ class BaseWanAttentionBlock(WanAttentionBlock):
ffn_dim,
ffn2_dim,
num_heads,
window_size=(-1, -1),
qk_norm=True,
cross_attn_norm=False,
eps=1e-6,
@@ -810,7 +786,7 @@ class BaseWanAttentionBlock(WanAttentionBlock):
attention_mode='sdpa',
rope_func="comfy"
):
super().__init__(cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, window_size, qk_norm, cross_attn_norm, eps, attention_mode, rope_func)
super().__init__(cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm, cross_attn_norm, eps, attention_mode, rope_func)
self.block_id = block_id
def forward(self, x, vace_hints=None, vace_context_scale=[1.0], **kwargs):
@@ -820,7 +796,7 @@ class BaseWanAttentionBlock(WanAttentionBlock):
if self.block_id is not None:
for i in range(len(vace_hints)):
x = x + vace_hints[i][self.block_id].to(x.device) * vace_context_scale[i]
x.add_(vace_hints[i][self.block_id].to(x.device), alpha=vace_context_scale[i])
return x
class Head(nn.Module):
@@ -891,7 +867,7 @@ class WanModel(ModelMixin, ConfigMixin):
"""
ignore_for_config = [
'patch_size', 'cross_attn_norm', 'qk_norm', 'text_dim', 'window_size'
'patch_size', 'cross_attn_norm', 'qk_norm', 'text_dim'
]
_no_split_modules = ['WanAttentionBlock']
@@ -911,7 +887,6 @@ class WanModel(ModelMixin, ConfigMixin):
out_dim=16,
num_heads=16,
num_layers=32,
window_size=(-1, -1),
qk_norm=True,
cross_attn_norm=True,
eps=1e-6,
@@ -955,8 +930,6 @@ class WanModel(ModelMixin, ConfigMixin):
Number of attention heads
num_layers (`int`, *optional*, defaults to 32):
Number of transformer blocks
window_size (`tuple`, *optional*, defaults to (-1, -1)):
Window size for local attention (-1 indicates global attention)
qk_norm (`bool`, *optional*, defaults to True):
Enable query/key normalization
cross_attn_norm (`bool`, *optional*, defaults to False):
@@ -982,7 +955,6 @@ class WanModel(ModelMixin, ConfigMixin):
self.out_dim = out_dim
self.num_heads = num_heads
self.num_layers = num_layers
self.window_size = window_size
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
@@ -1049,7 +1021,7 @@ class WanModel(ModelMixin, ConfigMixin):
# vace blocks
self.vace_blocks = nn.ModuleList([
VaceWanAttentionBlock('t2v_cross_attn', self.in_features, self.out_features, self.ffn_dim, self.ffn2_dim,self.num_heads, self.window_size, self.qk_norm,
VaceWanAttentionBlock('t2v_cross_attn', self.in_features, self.out_features, self.ffn_dim, self.ffn2_dim,self.num_heads, self.qk_norm,
self.cross_attn_norm, self.eps, block_id=i, attention_mode=self.attention_mode, rope_func=self.rope_func)
for i in self.vace_layers
])
@@ -1060,7 +1032,7 @@ class WanModel(ModelMixin, ConfigMixin):
)
self.blocks = nn.ModuleList([
BaseWanAttentionBlock('t2v_cross_attn', self.in_features, self.out_features, ffn_dim, self.ffn2_dim, num_heads,
window_size, qk_norm, cross_attn_norm, eps,
qk_norm, cross_attn_norm, eps,
attention_mode=self.attention_mode, rope_func=self.rope_func,
block_id=self.vace_layers_mapping[i] if i in self.vace_layers else None)
for i in range(num_layers)
@@ -1076,7 +1048,7 @@ class WanModel(ModelMixin, ConfigMixin):
self.blocks = nn.ModuleList([
WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads,
window_size, qk_norm, cross_attn_norm, eps,
qk_norm, cross_attn_norm, eps,
attention_mode=self.attention_mode, rope_func=self.rope_func)
for _ in range(num_layers)
])
@@ -1228,20 +1200,32 @@ class WanModel(ModelMixin, ConfigMixin):
if c.shape[1] > x.shape[1]:
c = c[:, :x.shape[1]]
c_list = [c]
hints = []
current_c = c
for b, block in enumerate(self.vace_blocks):
if b <= self.vace_blocks_to_swap and self.vace_blocks_to_swap >= 0:
block.to(self.main_device)
c_list = block(
c_list, x,
intermediate_device=self.offload_device if self.vace_blocks_to_swap != -1 else self.main_device,
nonblocking=self.use_non_blocking,
**kwargs)
if b == 0:
c_processed = block.before_proj(current_c) + x
else:
c_processed = current_c
c_processed = block.forward(c_processed, **kwargs)
# Store skip connection
c_skip = block.after_proj(c_processed)
hints.append(c_skip.to(
self.offload_device if self.vace_blocks_to_swap != -1 else self.main_device,
non_blocking=self.use_non_blocking
))
current_c = c_processed
if b <= self.vace_blocks_to_swap and self.vace_blocks_to_swap >= 0:
block.to(self.offload_device, non_blocking=self.use_non_blocking)
hints = c_list[:-1]
return hints
def forward(