refactor self attention

This commit is contained in:
kijai
2025-07-11 23:00:59 +03:00
parent 880194a1f9
commit 127bfa2c99
4 changed files with 95 additions and 111 deletions
+12 -7
View File
@@ -19,7 +19,7 @@ from .wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
from .multitalk.multitalk import timestep_transform, add_noise
from .enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight, set_num_frames
from .enhance_a_video.globals import set_enhance_weight, set_num_frames
from .taehv import TAEHV
from einops import rearrange
@@ -2308,6 +2308,12 @@ class WanVideoSampler:
rope_params(1024, 2 * (d // 6))
],
dim=1)
transformer.rope_func = rope_function
for block in transformer.blocks:
block.rope_func = rope_function
if transformer.vace_layers is not None:
for block in transformer.vace_blocks:
block.rope_func = rope_function
if not isinstance(cfg, list):
cfg = [cfg] * (steps +1)
@@ -2388,10 +2394,10 @@ class WanVideoSampler:
set_num_frames(context_frames)
else:
set_num_frames(latent_video_length)
enable_enhance()
enhance_enabled = True
else:
feta_args = None
disable_enhance()
enhance_enabled = False
# Initialize Cache if enabled
transformer.enable_teacache = transformer.enable_magcache = False
@@ -2628,6 +2634,7 @@ class WanVideoSampler:
't': timestep,
'current_step': idx,
'control_lora_enabled': control_lora_enabled,
'enhance_enabled': enhance_enabled,
'camera_embed': camera_embed,
'unianim_data': unianim_data,
'fun_ref': fun_ref_input if fun_ref_image is not None else None,
@@ -2642,7 +2649,6 @@ class WanVideoSampler:
"nag_context": text_embeds.get("nag_prompt_embeds", None),
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None,
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None,
"rope_func": rope_function
}
batch_size = 1
@@ -2877,10 +2883,9 @@ class WanVideoSampler:
latent_model_input = torch.cat([latent_model_input[:, shift_idx:]] + [latent_model_input[:, :shift_idx]], dim=1)
#enhance-a-video
enhance_enabled = False
if feta_args is not None and feta_start_percent <= current_step_percentage <= feta_end_percent:
enable_enhance()
else:
disable_enhance()
enhance_enabled = True
#flow-edit
if flowedit_args is not None:
+1
View File
@@ -668,6 +668,7 @@ class WanVideoModelLoader:
"num_heads": num_heads,
"num_layers": num_layers,
"attention_mode": attention_mode,
"rope_func": "comfy",
"main_device": device,
"offload_device": offload_device,
"teacache_coefficients": teacache_coefficients_map[model_variant],
+4
View File
@@ -33,6 +33,10 @@ def get_module_memory_mb(module):
memory += param.nelement() * param.element_size()
return memory / (1024 * 1024) # Convert to MB
def get_tensor_memory(tensor):
memory_bytes = tensor.element_size() * tensor.nelement()
return f"{memory_bytes / (1024 * 1024):.2f} MB"
def apply_lora(model, device_to, transformer_load_device, params_to_keep=None, dtype=None, base_dtype=None, state_dict=None, low_mem_load=False):
to_load = []
for n, m in model.model.named_modules():
+78 -104
View File
@@ -7,7 +7,6 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.modeling_utils import ModelMixin
from einops import repeat, rearrange
from ...enhance_a_video.enhance import get_feta_scores
from ...enhance_a_video.globals import is_enhance_enabled
try:
from torch.nn.attention.flex_attention import create_block_mask, flex_attention, BlockMask
@@ -234,8 +233,15 @@ class WanSelfAttention(nn.Module):
self.o = nn.Linear(in_features, out_features)
self.norm_q = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
def qkv_fn(self, x):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
q = self.norm_q(self.q(x)).view(b, s, n, d)
k = self.norm_k(self.k(x)).view(b, s, n, d)
v = self.v(x).view(b, s, n, d)
return q, k, v
def forward(self, x, seq_lens, grid_sizes, freqs, rope_func = "default", block_mask=None, ref_target_masks=None):
def forward(self, q, k, v, seq_lens, grid_sizes, freqs, block_mask=None):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -243,42 +249,23 @@ class WanSelfAttention(nn.Module):
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
# query, key, value
q = self.norm_q(self.q(x)).view(b, s, n, d)
k = self.norm_k(self.k(x)).view(b, s, n, d)
v = self.v(x).view(b, s, n, d)
if is_enhance_enabled():
feta_scores = get_feta_scores(q, k)
if self.attention_mode == 'flex_attention':
if rope_func == "comfy":
roped_query, roped_key = apply_rope_comfy(q, k, freqs)
else:
roped_query = rope_apply(q, grid_sizes, freqs).type_as(v)
roped_key = rope_apply(k, grid_sizes, freqs).type_as(v)
padded_length = math.ceil(q.shape[1] / 128) * 128 - q.shape[1]
padded_roped_query = torch.cat(
[roped_query,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
[q, torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)], dim=1
)
padded_roped_key = torch.cat(
[roped_key, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
[k, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)], dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
device=v.device, dtype=v.dtype)], dim=1
)
x = flex_attention(
query=padded_roped_query.transpose(2, 1),
@@ -286,39 +273,21 @@ class WanSelfAttention(nn.Module):
value=padded_v.transpose(2, 1),
block_mask=block_mask
)[:, :, :-padded_length].transpose(2, 1)
else:
if rope_func == "comfy":
q, k = apply_rope_comfy(q, k, freqs)
elif rope_func == "comfy_chunked":
q, k = apply_rope_comfy_chunked(q, k, freqs)
else:
q=rope_apply(q, grid_sizes, freqs)
k=rope_apply(k, grid_sizes, freqs)
else:
x = attention(
q=q,
k=k,
v=v,
q, k, v,
k_lens=seq_lens,
window_size=self.window_size,
attention_mode=self.attention_mode)
attention_mode=self.attention_mode
)
# output
x = x.flatten(2)
x = self.o(x)
if is_enhance_enabled():
x *= feta_scores
#multitalk
x_ref_attn_map = None
if ref_target_masks is not None:
x_ref_attn_map = get_attn_map_with_target(q.type_as(x), k.type_as(x), grid_sizes[0],
ref_target_masks=ref_target_masks)
return x, x_ref_attn_map
return x
def forward_split(self, x, seq_lens, grid_sizes, freqs, seq_chunks=1,current_step=0, video_attention_split_steps = [], rope_func = "default"):
def forward_split(self, q, k, v, seq_lens, grid_sizes, freqs, seq_chunks=1,current_step=0, video_attention_split_steps = []):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -326,25 +295,6 @@ class WanSelfAttention(nn.Module):
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
# query, key, value function
def qkv_fn(x):
q = self.norm_q(self.q(x)).view(b, s, n, d)
k = self.norm_k(self.k(x)).view(b, s, n, d)
v = self.v(x).view(b, s, n, d)
return q, k, v
q, k, v = qkv_fn(x)
if rope_func == "comfy":
q, k = apply_rope_comfy(q, k, freqs)
else:
q=rope_apply(q, grid_sizes, freqs)
k=rope_apply(k, grid_sizes, freqs)
if is_enhance_enabled():
feta_scores = get_feta_scores(q, k)
# Split by frames if multiple prompts are provided
if seq_chunks > 1 and current_step in video_attention_split_steps:
@@ -407,9 +357,6 @@ class WanSelfAttention(nn.Module):
x = x.flatten(2)
x = self.o(x)
if is_enhance_enabled():
x *= feta_scores
return x
def normalized_attention_guidance(self, b, n, d, q, context, nag_context=None, nag_params={}):
@@ -568,7 +515,9 @@ class WanAttentionBlock(nn.Module):
qk_norm=True,
cross_attn_norm=False,
eps=1e-6,
attention_mode='sdpa'):
attention_mode='sdpa',
rope_func="comfy"
):
super().__init__()
self.dim = out_features
self.ffn_dim = ffn_dim
@@ -578,6 +527,7 @@ class WanAttentionBlock(nn.Module):
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.attention_mode = attention_mode
self.rope_func = rope_func
# layers
self.norm1 = WanLayerNorm(out_features, eps)
@@ -602,7 +552,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
@@ -628,13 +578,13 @@ class WanAttentionBlock(nn.Module):
context_lens,
current_step,
video_attention_split_steps=[],
rope_func = "default",
clip_embed=None,
camera_embed=None,
audio_proj=None,
audio_context_lens=None,
audio_scale=1.0,
num_latent_frames=21,
enhance_enabled=False,
block_mask=None,
nag_params={},
nag_context=None,
@@ -665,28 +615,52 @@ class WanAttentionBlock(nn.Module):
input_x += camera_embed
# self-attention
x_ref_attn_map = None
#query, key, value
q, k, v = self.self_attn.qkv_fn(input_x)
del input_x
# FETA
if enhance_enabled:
feta_scores = get_feta_scores(q, k)
#RoPE
if self.rope_func == "comfy":
q, k = apply_rope_comfy(q, k, freqs)
elif self.rope_func == "comfy_chunked":
q, k = apply_rope_comfy_chunked(q, k, freqs)
else:
q=rope_apply(q, grid_sizes, freqs)
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:
y = self.self_attn.forward_split(
input_x,
seq_lens, grid_sizes,
freqs, rope_func=rope_func,
q, k, v,
seq_lens, grid_sizes, freqs,
seq_chunks=max(context.shape[0], clip_embed.shape[0] if clip_embed is not None else 0),
current_step=current_step,
video_attention_split_steps=video_attention_split_steps
)
else:
y, x_ref_attn_map = self.self_attn.forward(
input_x,
y = self.self_attn.forward(
q, k, v,
seq_lens, grid_sizes,
freqs, rope_func=rope_func,
block_mask=block_mask,
ref_target_masks=ref_target_masks,
freqs, block_mask=block_mask
)
#multitalk mask
if ref_target_masks is not None:
x_ref_attn_map = get_attn_map_with_target(q.type_as(x), k.type_as(x), grid_sizes[0], ref_target_masks=ref_target_masks)
# FETA
if enhance_enabled:
y.mul_(feta_scores)
#ReCamMaster
if camera_embed is not None:
y = self.projector(y)
del input_x
y = self.projector(y)
x.add_(y.mul_(e[2]))
del y
@@ -707,7 +681,6 @@ class WanAttentionBlock(nn.Module):
x.add_(y.mul_(e[5]))
del y
del e
return x
#@torch.compiler.disable()
@@ -796,9 +769,11 @@ class VaceWanAttentionBlock(WanAttentionBlock):
qk_norm=True,
cross_attn_norm=False,
eps=1e-6,
block_id=0
block_id=0,
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)
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)
self.block_id = block_id
if block_id == 0:
self.before_proj = nn.Linear(in_features, out_features)
@@ -832,9 +807,10 @@ class BaseWanAttentionBlock(WanAttentionBlock):
cross_attn_norm=False,
eps=1e-6,
block_id=None,
attention_mode='sdpa'
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)
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)
self.block_id = block_id
def forward(self, x, vace_hints=None, vace_context_scale=[1.0], **kwargs):
@@ -940,6 +916,7 @@ class WanModel(ModelMixin, ConfigMixin):
cross_attn_norm=True,
eps=1e-6,
attention_mode='sdpa',
rope_func='comfy',
main_device=torch.device('cuda'),
offload_device=torch.device('cpu'),
teacache_coefficients=[],
@@ -1010,6 +987,7 @@ class WanModel(ModelMixin, ConfigMixin):
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.attention_mode = attention_mode
self.rope_func = rope_func
self.main_device = main_device
self.offload_device = offload_device
@@ -1072,7 +1050,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,
self.cross_attn_norm, self.eps, block_id=i)
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
])
@@ -1083,7 +1061,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,
attention_mode=self.attention_mode,
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)
])
@@ -1099,7 +1077,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,
attention_mode=self.attention_mode)
attention_mode=self.attention_mode, rope_func=self.rope_func)
for _ in range(num_layers)
])
@@ -1280,6 +1258,7 @@ class WanModel(ModelMixin, ConfigMixin):
y=None,
device=torch.device('cuda'),
freqs=None,
enhance_enabled=False,
pred_id=None,
control_lora_enabled=False,
vace_data=None,
@@ -1298,8 +1277,7 @@ class WanModel(ModelMixin, ConfigMixin):
nag_params={},
nag_context=None,
multitalk_audio=None,
ref_target_masks=None,
rope_func="comfy",
ref_target_masks=None
):
r"""
Forward pass through the diffusion model
@@ -1642,13 +1620,13 @@ class WanModel(ModelMixin, ConfigMixin):
context=context,
context_lens=context_lens,
clip_embed=clip_embed,
rope_func=rope_func,
current_step=current_step,
video_attention_split_steps=self.video_attention_split_steps,
camera_embed=camera_embed,
audio_proj=audio_proj,
audio_context_lens=audio_context_lens,
num_latent_frames = F,
enhance_enabled=enhance_enabled,
audio_scale=audio_scale,
block_mask=self.block_mask,
nag_params=nag_params,
@@ -1837,7 +1815,3 @@ def relative_l1_distance(last_tensor, current_tensor):
norm = torch.abs(last_tensor).mean()
relative_l1_distance = l1_distance / norm
return relative_l1_distance.to(torch.float32).to(current_tensor.device)
def get_tensor_memory(tensor):
memory_bytes = tensor.element_size() * tensor.nelement()
return f"{memory_bytes / (1024 * 1024):.2f} MB"