Use proper self_attn output layers, and cleanup

This commit is contained in:
kijai
2025-11-04 21:13:29 +02:00
parent 977f4a5c3a
commit a51a53d5b7
2 changed files with 55 additions and 57 deletions
+22 -20
View File
@@ -485,28 +485,21 @@ class WanVideoSampler:
phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0)
# Video-as-prompt (VAP)
mot_ref_clip_embeds = mot_ref_context = None
mot_ref_clip_embeds = mot_ref_context = x_mot_ref = None
video_prompt_embeds = image_embeds.get("video_prompt_embeds", None)
if video_prompt_embeds is not None:
image_cond_mot_ref = video_prompt_embeds.get("image_embeds", None)
print("image_cond_mot_ref shape:", image_cond_mot_ref.shape)
image_cond_mask_ = video_prompt_embeds.get("mask", None)
if image_cond_mask_ is not None:
image_cond_mot_ref = torch.cat([image_cond_mask_, image_cond_mot_ref])
latents_mot_ref = video_prompt_embeds.get("video_prompt_latents", None)
print("latents_mot_ref shape:", latents_mot_ref.shape)
x_mot_ref = torch.cat([latents_mot_ref, image_cond_mot_ref], dim=0)
mot_ref_context = video_prompt_embeds.get("text_embeds", None)
mot_ref_clip_embeds = video_prompt_embeds.get("clip_context", None)
print("x_mot_ref shape:", x_mot_ref.shape)
# CLIP image features
clip_fea = image_embeds.get("clip_context", None)
if clip_fea is not None:
clip_fea = clip_fea.to(dtype)
clip_fea_neg = image_embeds.get("negative_clip_context", None)
if clip_fea_neg is not None:
clip_fea_neg = clip_fea_neg.to(dtype)
num_frames = image_embeds.get("num_frames", 0)
@@ -1057,7 +1050,7 @@ class WanVideoSampler:
if standin_input is not None:
rope_function = "comfy" # only works with this currently
freqs = None
freqs = freqs_mot_ref = None
transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None
d = transformer.dim // transformer.num_heads
@@ -1071,15 +1064,19 @@ class WanVideoSampler:
rope_params_mocha(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index, start=-1),
rope_params_mocha(1024, 2 * (d // 6), start=-1),
rope_params_mocha(1024, 2 * (d // 6), start=-1)
],
dim=1)
], dim=1)
elif "default" in rope_function or bidirectional_sampling: # original RoPE
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
],
dim=1)
], dim=1).to(device)
if x_mot_ref is not None:
freqs_mot_ref = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index, mot_ref_latent=x_mot_ref),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
], dim=1).to(device)
elif "comfy" in rope_function: # comfy's rope
transformer.rope_embedder.k = riflex_freq_index
transformer.rope_embedder.num_frames = latent_video_length
@@ -1364,6 +1361,10 @@ class WanVideoSampler:
z = z * c_in
timestep = c_noise
x_mot_ref_input = None
if image_cond_mot_ref is not None:
x_mot_ref_input = [x_mot_ref.to(z)]
base_params = {
'x': [z], # latent
'y': [image_cond_input] if image_cond_input is not None else None, # image cond
@@ -1420,9 +1421,10 @@ class WanVideoSampler:
"flashvsr_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling
"flashvsr_strength": flashvsr_strength, # FlashVSR strength
"num_cond_latents": len(all_indices) if transformer.is_longcat else None, # number of cond latents LongCat to separate attention
"x_mot_ref": [x_mot_ref.to(z)] if image_cond_mot_ref is not None else None, # motion reference latents for VAP
"x_mot_ref": x_mot_ref_input, # motion reference latents for VAP
"mot_ref_context": mot_ref_context if image_cond_mot_ref is not None else None, # motion reference context for VAP
"mot_ref_clip_embeds": mot_ref_clip_embeds, # motion reference clip features for VAP
"freqs_mot_ref": freqs_mot_ref, # motion reference RoPE freqs for VAP
}
batch_size = 1
@@ -1577,23 +1579,23 @@ class WanVideoSampler:
noise_pred_uncond.view(batch_size, -1)
).view(batch_size, 1, 1, 1)
noise_pred_uncond_scaled = noise_pred_uncond * alpha
noise_pred_uncond = noise_pred_uncond * alpha
if use_tangential:
noise_pred_uncond_scaled = tangential_projection(noise_pred_cond, noise_pred_uncond_scaled)
noise_pred_uncond = tangential_projection(noise_pred_cond, noise_pred_uncond)
# RAAG (RATIO-aware Adaptive Guidance)
if raag_alpha > 0.0:
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond_scaled, cfg_scale, raag_alpha)
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond, cfg_scale, raag_alpha)
log.info(f"RAAG modified cfg: {cfg_scale}")
#https://github.com/WikiChao/FreSca
if use_fresca:
filtered_cond = fourier_filter(noise_pred_cond - noise_pred_uncond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff)
noise_pred = noise_pred_uncond_scaled + cfg_scale * filtered_cond * alpha
noise_pred = noise_pred_uncond + cfg_scale * filtered_cond * alpha
else:
noise_pred = noise_pred_uncond_scaled + cfg_scale * (noise_pred_cond - noise_pred_uncond_scaled)
del noise_pred_uncond_scaled, noise_pred_cond, noise_pred_uncond
noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond)
del noise_pred_uncond, noise_pred_cond
if latent_model_input_ovi is not None:
if ovi_audio_cfg is None:
+33 -37
View File
@@ -249,7 +249,7 @@ def sinusoidal_embedding_1d(dim, position):
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0, freqs_scaling=1.0):
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0, freqs_scaling=1.0, mot_ref_latent=None):
assert dim % 2 == 0
exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim)
inv_theta_pow = 1.0 / torch.pow(theta, exponents)
@@ -259,9 +259,16 @@ def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0, freqs_scaling=1.0
inv_theta_pow[k-1] = 0.9 * 2 * torch.pi / L_test
inv_theta_pow *= freqs_scaling
if mot_ref_latent is not None:
freqs = torch.arange(-mot_ref_latent.shape[1], max_seq_len)
freqs = torch.outer(freqs, inv_theta_pow)
freqs = torch.polar(torch.ones_like(freqs), freqs)
freqs = freqs[:max_seq_len]
freqs = torch.outer(torch.arange(max_seq_len), inv_theta_pow)
freqs = torch.polar(torch.ones_like(freqs), freqs)
else:
freqs = torch.outer(torch.arange(max_seq_len), inv_theta_pow)
freqs = torch.polar(torch.ones_like(freqs), freqs)
return freqs
@torch.autocast(device_type=mm.get_autocast_device(mm.get_torch_device()), enabled=False)
@@ -748,7 +755,7 @@ class WanT2VCrossAttention(WanSelfAttention):
return self.o(x)
class WanT2VCrossAttentionMOTRef(WanSelfAttention):
class WanCrossAttentionMOTRef(WanSelfAttention):
def __init__(self, in_features, out_features, num_heads, kv_dim=None, qk_norm=True, eps=1e-6, attention_mode='sdpa', rms_norm_function="default", head_norm=False):
super().__init__(in_features, out_features, num_heads, qk_norm, eps, kv_dim=kv_dim, rms_norm_function=rms_norm_function, head_norm=head_norm)
@@ -756,8 +763,6 @@ class WanT2VCrossAttentionMOTRef(WanSelfAttention):
self.v_img = nn.Linear(in_features, out_features)
self.norm_k_img = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
self.ip_adapter = None
self.k_fusion = None
def forward(self, x, context, grid_sizes=None, clip_embed=None, rope_func="comfy", **kwargs):
b, n, d = x.size(0), self.num_heads, self.head_dim
@@ -987,7 +992,7 @@ class WanAttentionBlock(nn.Module):
self.norm2_mot_ref = WanLayerNorm(self.dim, eps)
self.norm3_mot_ref = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
self.self_attn_mot_ref = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function)
self.cross_attn_mot_ref = WanT2VCrossAttentionMOTRef(in_features, out_features, num_heads, qk_norm, eps, rms_norm_function=rms_norm_function)
self.cross_attn_mot_ref = WanCrossAttentionMOTRef(in_features, out_features, num_heads, qk_norm, eps, rms_norm_function=rms_norm_function)
self.modulation_mot_ref = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
self.ffn_mot_ref = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features))
@@ -1124,11 +1129,6 @@ class WanAttentionBlock(nn.Module):
# video-as-prompt motion reference
use_mot_ref = x_mot_ref is not None and self.mot_ref_block
if use_mot_ref:
#import einops
# shift_msa_mot_ref, scale_msa_mot_ref, gate_msa_mot_ref, shift_mlp_mot_ref, scale_mlp_mot_ref, gate_mlp_mot_ref = self.get_mod(e_mot_ref.to(x.device), self.modulation_mot_ref)
# norm_x_mot_ref = einops.rearrange(self.norm1_mot_ref(x_mot_ref.to(scale_msa_mot_ref.dtype)), 'b (n t) c -> b n t c', n=num_mot_ref)
# input_x_mot_ref = self.modulate(norm_x_mot_ref, shift_msa_mot_ref, scale_msa_mot_ref).to(input_dtype)
# input_x_mot_ref = einops.rearrange(input_x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref)
shift_msa_mot_ref, scale_msa_mot_ref, gate_msa_mot_ref, shift_mlp_mot_ref, scale_mlp_mot_ref, gate_mlp_mot_ref = self.get_mod(e_mot_ref.to(x.device), self.modulation_mot_ref)
input_x_mot_ref = self.modulate(self.norm1_mot_ref(x_mot_ref.to(shift_msa_mot_ref.dtype)), shift_msa_mot_ref, scale_msa_mot_ref).to(input_dtype)
@@ -1192,6 +1192,9 @@ class WanAttentionBlock(nn.Module):
else:
q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time)
k = rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time)
if use_mot_ref:
q_mot_ref = rope_apply(q_mot_ref, grid_sizes_mot_ref, freqs_mot_ref)
k_mot_ref = rope_apply(k_mot_ref, grid_sizes_mot_ref, freqs_mot_ref)
if x_ovi is not None:
q_ovi, k_ovi, v_ovi = self.audio_block.self_attn.qkv_fn(input_x_ovi)
@@ -1253,13 +1256,19 @@ class WanAttentionBlock(nn.Module):
# merge x_cond and x_noise
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
elif use_mot_ref:
y_temp = self.self_attn_mot_ref.forward(
y_temp = attention(
torch.cat([q, q_mot_ref], dim=1),
torch.cat([k, k_mot_ref], dim=1),
torch.cat([v, v_mot_ref], dim=1),
seq_lens
attention_mode=self.attention_mode
)
y, y_mot_ref = torch.split(y_temp, [q.shape[1], q_mot_ref.shape[1]], dim=1)
y, y_mot_ref = (
y_temp[:, :q.shape[1]],
y_temp[:, q.shape[1]:q.shape[1]+q_mot_ref.shape[1]]
)
y = self.self_attn.o(y.flatten(2))
y_mot_ref = self.self_attn_mot_ref.o(y_mot_ref.flatten(2))
del y_temp
else:
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale)
@@ -1292,8 +1301,6 @@ class WanAttentionBlock(nn.Module):
if not is_longcat:
x = x.addcmul(y, gate_msa)
if use_mot_ref:
#x_mot_ref = x_mot_ref.addcmul(einops.rearrange(y_mot_ref, 'b (n t) c -> b n t c', n=num_mot_ref), gate_msa_mot_ref)
#x_mot_ref = einops.rearrange(x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref)
x_mot_ref = x_mot_ref.addcmul(y_mot_ref, gate_msa_mot_ref)
else:
x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C)
@@ -1392,13 +1399,6 @@ class WanAttentionBlock(nn.Module):
x_ip = x_ip.addcmul(y_ip, gate_mlp_ip)
if use_mot_ref:
# norm2_x_mot_ref = einops.rearrange(self.norm2_mot_ref(x_mot_ref.to(shift_mlp_mot_ref.dtype)), 'b (n t) c -> b n t c', n=num_mot_ref)
# mod_x_mot_ref = torch.addcmul(shift_mlp_mot_ref, norm2_x_mot_ref, 1 + scale_mlp_mot_ref)
# mod_x_mot_ref = einops.rearrange(mod_x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref)
# x_ffn_mot_ref = self.ffn_mot_ref(mod_x_mot_ref.to(input_dtype))
# x_ffn_mot_ref = einops.rearrange(x_ffn_mot_ref, 'b (n t) c -> b n t c', n=num_mot_ref)
# x_mot_ref = x_mot_ref.addcmul(x_ffn_mot_ref, gate_mlp_mot_ref)
# x_mot_ref = einops.rearrange(x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref)
norm2_x_mot_ref = self.norm2_mot_ref(x_mot_ref.to(shift_mlp_mot_ref.dtype))
mod_x_mot_ref = torch.addcmul(shift_mlp_mot_ref, norm2_x_mot_ref, 1 + scale_mlp_mot_ref)
x_ffn_mot_ref = self.ffn_mot_ref(mod_x_mot_ref.to(input_dtype))
@@ -1603,10 +1603,10 @@ class MLPProj(torch.nn.Module):
if fl_pos_emb: # NOTE: we only use this for `fl2v`
self.emb_pos = nn.Parameter(torch.zeros(1, 257 * 2, 1280))
def forward(self, image_embeds):
def forward(self, image_embeds, dtype=torch.float32):
if hasattr(self, 'emb_pos'):
image_embeds = image_embeds + self.emb_pos.to(image_embeds.device)
clip_extra_context_tokens = self.proj(image_embeds.to(self.proj[1].weight.dtype)).to(image_embeds.dtype)
clip_extra_context_tokens = self.proj(image_embeds.to(self.proj[1].weight.dtype)).to(dtype)
return clip_extra_context_tokens
from .s2v.auxi_blocks import MotionEncoder_tc
@@ -1870,9 +1870,6 @@ class WanModel(torch.nn.Module):
ConvMLP(dim, dim * 4, kernel_size=7, padding=3),
)
if is_VAP:
self.patch_embedding_mot_ref = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.original_patch_embedding = self.patch_embedding
self.expanded_patch_embedding = self.patch_embedding
@@ -1889,12 +1886,13 @@ class WanModel(torch.nn.Module):
adaln_tembed_dim = 512
self.time_embedding = TimestepEmbedder(t_embed_dim=adaln_tembed_dim, frequency_embedding_size=freq_dim)
if is_VAP:
self.patch_embedding_mot_ref = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.time_embedding_mot_ref = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
self.time_projection_mot_ref = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
self.text_embedding_mot_ref = nn.Sequential(nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'), nn.Linear(dim, dim))
self.img_emb_mot_ref = MLPProj(1280, dim)
VAP_layers = [0, 4, 8, 12, 16, 20, 24, 28, 32, 36]
if vace_layers is not None:
self.vace_layers = [i for i in range(0, self.num_layers, 2)] if vace_layers is None else vace_layers
@@ -1929,15 +1927,13 @@ class WanModel(torch.nn.Module):
else:
cross_attn_type = 'no_cross_attn'
VAP_layers = [0, 4, 8, 12, 16, 20, 24, 28, 32, 36]
self.blocks = nn.ModuleList([
WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm, cross_attn_norm, eps,
attention_mode=self.attention_mode, rope_func=self.rope_func, rms_norm_function=rms_norm_function,
use_motion_attn=(i % 4 == 0 and use_motion_attn), use_humo_audio_attn=self.humo_audio,
face_fuser_block = (i % 5 == 0 and is_wananimate), lynx_ip_layers=lynx_ip_layers, lynx_ref_layers=lynx_ref_layers,
block_idx=i, is_longcat=is_longcat, mot_ref_block=i in VAP_layers and is_VAP)
block_idx=i, is_longcat=is_longcat, mot_ref_block=is_VAP and i in VAP_layers)
for i in range(num_layers)
])
#MTV Crafter
@@ -2211,13 +2207,13 @@ class WanModel(torch.nn.Module):
return x.add(residual_out, alpha=strength)
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None, mot=False):
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None, mot_ref=False):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if mot:
if mot_ref:
t_start = -t_len
if steps_t is None:
@@ -2301,7 +2297,7 @@ class WanModel(torch.nn.Module):
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
num_cond_latents=None,
x_mot_ref=None, mot_ref_context=None, mot_ref_clip_embeds=None,
x_mot_ref=None, mot_ref_context=None, mot_ref_clip_embeds=None, freqs_mot_ref=None,
):
r"""
Forward pass through the diffusion model
@@ -2552,7 +2548,7 @@ class WanModel(torch.nn.Module):
self.cached_ntk_alphas = ntk_alphas
if x_mot_ref is not None:
freqs_mot_ref = self.rope_encode_comfy(F, H, W, mot=True, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype)
freqs_mot_ref = self.rope_encode_comfy(F, H, W, mot_ref=True, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype)
# Stand-In RoPE frequencies