7 Commits
Author SHA1 Message Date
kijai 7a0da7708e Merge branch 'main' into vap 2025-11-04 23:15:20 +02:00
kijai a51a53d5b7 Use proper self_attn output layers, and cleanup 2025-11-04 21:13:29 +02:00
kijai 977f4a5c3a Update model.py 2025-11-04 11:31:08 +02:00
kijai 2a45675498 Merge branch 'main' into vap 2025-11-04 10:39:35 +02:00
kijai ea414c54ac Create wanvideo_I2V_video-as-prompt_testing_WIP.json 2025-11-01 16:53:15 +02:00
kijai 0013ae0ece Init VAP 2025-11-01 16:47:37 +02:00
kijai 0e904e6035 Remove unnecessary casts 2025-10-31 17:28:12 +02:00
5 changed files with 2616 additions and 85 deletions
File diff suppressed because it is too large Load Diff
+25
View File
@@ -22,6 +22,30 @@ offload_device = mm.unet_offload_device()
VAE_STRIDE = (4, 8, 8) VAE_STRIDE = (4, 8, 8)
PATCH_SIZE = (1, 2, 2) PATCH_SIZE = (1, 2, 2)
class WanVideoAddVideoPromptEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"image_embeds": ("WANVIDIMAGE_EMBEDS",),
"video_prompt_embeds": ("WANVIDIMAGE_EMBEDS",),
"video_prompt_latents": ("LATENT", ),
"text_embeds": ("WANVIDEOTEXTEMBEDS", ),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
EXPERIMENTAL = True
def add(self, image_embeds, video_prompt_embeds, video_prompt_latents, text_embeds):
updated = dict(image_embeds)
updated["video_prompt_embeds"] = video_prompt_embeds
updated["video_prompt_embeds"]["video_prompt_latents"] = video_prompt_latents["samples"][0]
updated["video_prompt_embeds"]["text_embeds"] = text_embeds
return (updated,)
class WanVideoEnhanceAVideo: class WanVideoEnhanceAVideo:
@classmethod @classmethod
@@ -2206,6 +2230,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds, "WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents, "WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE, "WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
"WanVideoAddVideoPromptEmbeds": WanVideoAddVideoPromptEmbeds,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
+1
View File
@@ -1349,6 +1349,7 @@ class WanVideoModelLoader:
"lynx_ip_layers": lynx_ip_layers, "lynx_ip_layers": lynx_ip_layers,
"lynx_ref_layers": lynx_ref_layers, "lynx_ref_layers": lynx_ref_layers,
"is_longcat": dim == 4096, "is_longcat": dim == 4096,
"is_VAP": True if "patch_embedding_mot_ref.weight" in sd else False
} }
+38 -16
View File
@@ -296,6 +296,7 @@ class WanVideoSampler:
phantom_latents = fun_ref_image = ATI_tracks = None phantom_latents = fun_ref_image = ATI_tracks = None
add_cond = attn_cond = attn_cond_neg = noise_pred_flipped = None add_cond = attn_cond = attn_cond_neg = noise_pred_flipped = None
humo_audio = humo_audio_neg = None humo_audio = humo_audio_neg = None
image_cond_mot_ref = None
#I2V #I2V
image_cond = image_embeds.get("image_embeds", None) image_cond = image_embeds.get("image_embeds", None)
@@ -483,13 +484,22 @@ class WanVideoSampler:
phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0) phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0)
phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0) phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0)
# Video-as-prompt (VAP)
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)
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)
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)
# CLIP image features # CLIP image features
clip_fea = image_embeds.get("clip_context", None) 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) 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) num_frames = image_embeds.get("num_frames", 0)
@@ -1040,7 +1050,7 @@ class WanVideoSampler:
if standin_input is not None: if standin_input is not None:
rope_function = "comfy" # only works with this currently rope_function = "comfy" # only works with this currently
freqs = None freqs = freqs_mot_ref = None
transformer.rope_embedder.k = None transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None transformer.rope_embedder.num_frames = None
d = transformer.dim // transformer.num_heads d = transformer.dim // transformer.num_heads
@@ -1054,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, 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),
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 elif "default" in rope_function or bidirectional_sampling: # original RoPE
freqs = torch.cat([ freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index), 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)),
rope_params(1024, 2 * (d // 6)) rope_params(1024, 2 * (d // 6))
], ], dim=1).to(device)
dim=1) 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 elif "comfy" in rope_function: # comfy's rope
transformer.rope_embedder.k = riflex_freq_index transformer.rope_embedder.k = riflex_freq_index
transformer.rope_embedder.num_frames = latent_video_length transformer.rope_embedder.num_frames = latent_video_length
@@ -1347,6 +1361,10 @@ class WanVideoSampler:
z = z * c_in z = z * c_in
timestep = c_noise 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 = { base_params = {
'x': [z], # latent 'x': [z], # latent
'y': [image_cond_input] if image_cond_input is not None else None, # image cond 'y': [image_cond_input] if image_cond_input is not None else None, # image cond
@@ -1402,7 +1420,11 @@ class WanVideoSampler:
"ovi_negative_text_embeds": ovi_negative_text_embeds, # Audio latent model negative text embeds for Ovi "ovi_negative_text_embeds": ovi_negative_text_embeds, # Audio latent model negative text embeds for Ovi
"flashvsr_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling "flashvsr_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling
"flashvsr_strength": flashvsr_strength, # FlashVSR strength "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 "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_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 batch_size = 1
@@ -1557,23 +1579,23 @@ class WanVideoSampler:
noise_pred_uncond.view(batch_size, -1) noise_pred_uncond.view(batch_size, -1)
).view(batch_size, 1, 1, 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: 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) # RAAG (RATIO-aware Adaptive Guidance)
if raag_alpha > 0.0: 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}") log.info(f"RAAG modified cfg: {cfg_scale}")
#https://github.com/WikiChao/FreSca #https://github.com/WikiChao/FreSca
if use_fresca: if use_fresca:
filtered_cond = fourier_filter(noise_pred_cond - noise_pred_uncond, fresca_scale_low, fresca_scale_high, fresca_freq_cutoff) 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: else:
noise_pred = noise_pred_uncond_scaled + cfg_scale * (noise_pred_cond - noise_pred_uncond_scaled) noise_pred = noise_pred_uncond + cfg_scale * (noise_pred_cond - noise_pred_uncond)
del noise_pred_uncond_scaled, noise_pred_cond, noise_pred_uncond del noise_pred_uncond, noise_pred_cond
if latent_model_input_ovi is not None: if latent_model_input_ovi is not None:
if ovi_audio_cfg is None: if ovi_audio_cfg is None:
+182 -69
View File
@@ -249,7 +249,7 @@ def sinusoidal_embedding_1d(dim, position):
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1) x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x 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 assert dim % 2 == 0
exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim) exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim)
inv_theta_pow = 1.0 / torch.pow(theta, exponents) 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[k-1] = 0.9 * 2 * torch.pi / L_test
inv_theta_pow *= freqs_scaling 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) else:
freqs = torch.polar(torch.ones_like(freqs), freqs) freqs = torch.outer(torch.arange(max_seq_len), inv_theta_pow)
freqs = torch.polar(torch.ones_like(freqs), freqs)
return freqs return freqs
@torch.autocast(device_type=mm.get_autocast_device(mm.get_torch_device()), enabled=False) @torch.autocast(device_type=mm.get_autocast_device(mm.get_torch_device()), enabled=False)
@@ -356,10 +363,10 @@ class WanRMSNorm(nn.Module):
if use_chunked: if use_chunked:
return self.forward_chunked(x, num_chunks) return self.forward_chunked(x, num_chunks)
else: else:
return self._norm(x.to(self.weight.dtype)) * self.weight return (self._norm(x.to(self.weight.dtype)) * self.weight).to(x.dtype)
def _norm(self, x): def _norm(self, x):
return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)).to(x.dtype) return x * (torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps))
def forward_chunked(self, x, num_chunks=4): def forward_chunked(self, x, num_chunks=4):
output = torch.empty_like(x) output = torch.empty_like(x)
@@ -386,7 +393,7 @@ class WanFusedRMSNorm(nn.RMSNorm):
if use_chunked: if use_chunked:
return self.forward_chunked(x, num_chunks) return self.forward_chunked(x, num_chunks)
else: else:
return super().forward(x) return super().forward(x.to(self.weight.dtype).to(x.dtype))
def forward_chunked(self, x, num_chunks=4): def forward_chunked(self, x, num_chunks=4):
output = torch.empty_like(x) output = torch.empty_like(x)
@@ -398,7 +405,7 @@ class WanFusedRMSNorm(nn.RMSNorm):
for size in chunk_sizes: for size in chunk_sizes:
end_idx = start_idx + size end_idx = start_idx + size
chunk = x[:, start_idx:end_idx, :] chunk = x[:, start_idx:end_idx, :]
output[:, start_idx:end_idx, :] = super().forward(chunk) output[:, start_idx:end_idx, :] = super().forward(chunk.to(self.weight.dtype)).to(chunk.dtype)
start_idx = end_idx start_idx = end_idx
return output return output
@@ -464,8 +471,8 @@ class WanSelfAttention(nn.Module):
def qkv_fn(self, x): def qkv_fn(self, x):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d) q = self.norm_q(self.q(x)).view(b, s, n, d)
k = self.norm_k(self.k(x).to(self.norm_k.weight.dtype)).to(x.dtype).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) v = self.v(x).view(b, s, n, d)
return q, k, v return q, k, v
@@ -480,8 +487,8 @@ class WanSelfAttention(nn.Module):
def qkv_fn_ip(self, x): def qkv_fn_ip(self, x):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
q = self.norm_q(self.q(x) + self.q_loras(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d) q = self.norm_q(self.q(x) + self.q_loras(x)).view(b, s, n, d)
k = self.norm_k(self.k(x) + self.k_loras(x).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d) k = self.norm_k(self.k(x) + self.k_loras(x)).view(b, s, n, d)
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d) v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
return q, k, v return q, k, v
@@ -674,7 +681,7 @@ class WanT2VCrossAttention(WanSelfAttention):
if num_cond_latents is not None and num_cond_latents > 0: if num_cond_latents is not None and num_cond_latents > 0:
num_cond_latents_thw = num_cond_latents * (s // num_latent_frames) num_cond_latents_thw = num_cond_latents * (s // num_latent_frames)
x = x[:, num_cond_latents_thw:] x = x[:, num_cond_latents_thw:]
q = self.norm_q(self.q(x).view(b, -1, n, d)) q = self.norm_q(self.q(x).view(b, -1, n, d).to(self.norm_q.weight.dtype)).to(x.dtype)
else: else:
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d) q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d)
@@ -747,6 +754,32 @@ class WanT2VCrossAttention(WanSelfAttention):
return torch.cat([torch.zeros((b, num_cond_latents_thw, x.shape[-1]), dtype=x.dtype, device=x.device), self.o(x)], dim=1).contiguous() return torch.cat([torch.zeros((b, num_cond_latents_thw, x.shape[-1]), dtype=x.dtype, device=x.device), self.o(x)], dim=1).contiguous()
return self.o(x) return self.o(x)
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)
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()
self.attention_mode = attention_mode
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
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype), num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d)
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
x = attention(q, k, v, attention_mode=self.attention_mode).flatten(2)
if clip_embed is not None:
k_img = self.norm_k_img(self.k_img(clip_embed)).view(b, -1, n, d)
v_img = self.v_img(clip_embed).view(b, -1, n, d)
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2)
x = x + img_x
return self.o(x)
class WanI2VCrossAttention(WanSelfAttention): class WanI2VCrossAttention(WanSelfAttention):
@@ -773,13 +806,13 @@ class WanI2VCrossAttention(WanSelfAttention):
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
else: else:
# text attention # text attention
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype) k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d) v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, attention_mode=self.attention_mode).flatten(2) x_text = attention(q, k, v, attention_mode=self.attention_mode).flatten(2)
#img attention #img attention
if clip_embed is not None: if clip_embed is not None:
k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype) k_img = self.norm_k_img(self.k_img(clip_embed)).view(b, -1, n, d)
v_img = self.v_img(clip_embed).view(b, -1, n, d) v_img = self.v_img(clip_embed).view(b, -1, n, d)
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2) img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2)
x = x_text + img_x x = x_text + img_x
@@ -871,8 +904,8 @@ class MTVCrafterMotionAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value # compute query, key, value
q = self.norm_q(self.q(x)).view(b, -1, n, d) q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, -1, n, d)
k = self.norm_k(self.k(mo)).view(b, n, -1, d) k = self.norm_k(self.k(mo).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, n, -1, d)
v = self.v(mo).view(b, -1, n, d) v = self.v(mo).view(b, -1, n, d)
# compute attention # compute attention
@@ -897,7 +930,7 @@ class WanAttentionBlock(nn.Module):
cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default", qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default",
use_motion_attn=False, use_humo_audio_attn=False, face_fuser_block=False, lynx_ip_layers=None, lynx_ref_layers=None, use_motion_attn=False, use_humo_audio_attn=False, face_fuser_block=False, lynx_ip_layers=None, lynx_ref_layers=None,
block_idx=0, is_longcat=False): block_idx=0, mot_ref_block=False, is_longcat=False):
super().__init__() super().__init__()
self.dim = out_features self.dim = out_features
self.ffn_dim = ffn_dim self.ffn_dim = ffn_dim
@@ -913,6 +946,7 @@ class WanAttentionBlock(nn.Module):
self.dense_block = False self.dense_block = False
self.dense_attention_mode = "sageattn" self.dense_attention_mode = "sageattn"
self.block_idx = block_idx self.block_idx = block_idx
self.mot_ref_block = mot_ref_block
self.kv_cache = None self.kv_cache = None
self.use_motion_attn = use_motion_attn self.use_motion_attn = use_motion_attn
@@ -952,6 +986,16 @@ class WanAttentionBlock(nn.Module):
self.seg_idx = None self.seg_idx = None
# video-as-prompt (VAP)
if mot_ref_block:
self.norm1_mot_ref = WanLayerNorm(self.dim, eps)
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 = 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))
# HuMo audio cross-attn # HuMo audio cross-attn
if use_humo_audio_attn: if use_humo_audio_attn:
self.audio_cross_attn_wrapper = AudioCrossAttentionWrapper(in_features, out_features, num_heads, qk_norm, eps, kv_dim=1536) self.audio_cross_attn_wrapper = AudioCrossAttentionWrapper(in_features, out_features, num_heads, qk_norm, eps, kv_dim=1536)
@@ -1045,7 +1089,8 @@ class WanAttentionBlock(nn.Module):
lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx
x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None, x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None,
num_cond_latents=None, #longcat image cond amount num_cond_latents=None, #longcat image cond amount
): # VAP
x_mot_ref=None, context_mot_ref=None, grid_sizes_mot_ref=None, e_mot_ref=None, freqs_mot_ref=None, clip_embed_mot_ref=None, num_mot_ref=1):
r""" r"""
Args: Args:
x(Tensor): Shape [B, L, C] x(Tensor): Shape [B, L, C]
@@ -1081,6 +1126,12 @@ class WanAttentionBlock(nn.Module):
input_x = torch.concat([input_x, input_x_ip], dim=1) input_x = torch.concat([input_x, input_x_ip], dim=1)
self.kv_cache = None self.kv_cache = None
# video-as-prompt motion reference
use_mot_ref = x_mot_ref is not None and self.mot_ref_block
if use_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)
if x_ovi is not None: if x_ovi is not None:
shift_msa_ovi, scale_msa_ovi, gate_msa_ovi, shift_mlp_ovi, scale_mlp_ovi, gate_mlp_ovi = self.get_mod(e_ovi.to(x.device), self.audio_block.modulation) shift_msa_ovi, scale_msa_ovi, gate_msa_ovi, shift_mlp_ovi, scale_mlp_ovi, gate_mlp_ovi = self.get_mod(e_ovi.to(x.device), self.audio_block.modulation)
input_x_ovi = self.modulate(self.audio_block.norm1(x_ovi), shift_msa_ovi, scale_msa_ovi) input_x_ovi = self.modulate(self.audio_block.norm1(x_ovi), shift_msa_ovi, scale_msa_ovi)
@@ -1125,8 +1176,13 @@ class WanAttentionBlock(nn.Module):
q, k, v = self.self_attn.qkv_fn_longcat(input_x) q, k, v = self.self_attn.qkv_fn_longcat(input_x)
else: else:
q, k, v = self.self_attn.qkv_fn(input_x) q, k, v = self.self_attn.qkv_fn(input_x)
if use_mot_ref:
q_mot_ref, k_mot_ref, v_mot_ref = self.self_attn_mot_ref.qkv_fn(input_x_mot_ref)
# Apply RoPE
if self.rope_func == "comfy": if self.rope_func == "comfy":
q, k = apply_rope_comfy(q, k, freqs) q, k = apply_rope_comfy(q, k, freqs)
if use_mot_ref:
q_mot_ref, k_mot_ref = apply_rope_comfy(q_mot_ref, k_mot_ref, freqs_mot_ref)
elif self.rope_func == "comfy_chunked": elif self.rope_func == "comfy_chunked":
q, k = apply_rope_comfy_chunked(q, k, freqs) q, k = apply_rope_comfy_chunked(q, k, freqs)
elif self.rope_func == "mocha": elif self.rope_func == "mocha":
@@ -1136,6 +1192,9 @@ class WanAttentionBlock(nn.Module):
else: else:
q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time) q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time)
k = rope_apply(k, 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: if x_ovi is not None:
q_ovi, k_ovi, v_ovi = self.audio_block.self_attn.qkv_fn(input_x_ovi) q_ovi, k_ovi, v_ovi = self.audio_block.self_attn.qkv_fn(input_x_ovi)
@@ -1196,6 +1255,20 @@ class WanAttentionBlock(nn.Module):
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens) x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
# merge x_cond and x_noise # merge x_cond and x_noise
y = torch.cat([x_cond, x_noise], dim=1).contiguous() y = torch.cat([x_cond, x_noise], dim=1).contiguous()
elif use_mot_ref:
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),
attention_mode=self.attention_mode
)
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: else:
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale) y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale)
@@ -1227,8 +1300,11 @@ class WanAttentionBlock(nn.Module):
else: else:
if not is_longcat: if not is_longcat:
x = x.addcmul(y, gate_msa) x = x.addcmul(y, gate_msa)
if use_mot_ref:
x_mot_ref = x_mot_ref.addcmul(y_mot_ref, gate_msa_mot_ref)
else: else:
x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C) x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C)
del y, gate_msa del y, gate_msa
# cross-attention & ffn function # cross-attention & ffn function
@@ -1263,6 +1339,10 @@ class WanAttentionBlock(nn.Module):
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs, rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents) adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents)
x = x.to(input_dtype) x = x.to(input_dtype)
if use_mot_ref:
x_mot_ref = x_mot_ref + self.cross_attn_mot_ref(self.norm3_mot_ref(x_mot_ref.to(self.norm3_mot_ref.weight.dtype)).to(input_dtype), context_mot_ref, grid_sizes_mot_ref,
clip_embed=clip_embed_mot_ref)
x_mot_ref = x_mot_ref.to(input_dtype)
# MultiTalk # MultiTalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock): if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding, x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
@@ -1270,8 +1350,8 @@ class WanAttentionBlock(nn.Module):
x = x.add(x_audio, alpha=audio_scale) x = x.add(x_audio, alpha=audio_scale)
# MTV-Crafter Motion Attention # MTV-Crafter Motion Attention
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None: if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs) x_motion = self.motion_attn(self.norm4(x.to(self.norm4.weight.dtype)).to(input_dtype), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
x = x.add(x_motion, alpha=mtv_strength) x = x.add(x_motion, alpha=mtv_strength)
# HuMo Audio Cross-Attention # HuMo Audio Cross-Attention
@@ -1317,7 +1397,13 @@ class WanAttentionBlock(nn.Module):
x_ip = x_ip.addcmul(y_ip, gate_msa_ip) x_ip = x_ip.addcmul(y_ip, gate_msa_ip)
y_ip = self.ffn(torch.addcmul(shift_mlp_ip, self.norm2(x_ip), 1 + scale_mlp_ip)) y_ip = self.ffn(torch.addcmul(shift_mlp_ip, self.norm2(x_ip), 1 + scale_mlp_ip))
x_ip = x_ip.addcmul(y_ip, gate_mlp_ip) x_ip = x_ip.addcmul(y_ip, gate_mlp_ip)
return x, x_ip, lynx_ref_feature, x_ovi
if use_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))
x_mot_ref = x_mot_ref.addcmul(x_ffn_mot_ref, gate_mlp_mot_ref)
return x, x_ip, lynx_ref_feature, x_ovi, x_mot_ref
@torch.compiler.disable() @torch.compiler.disable()
def split_cross_attn_ffn(self, x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None): def split_cross_attn_ffn(self, x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None):
@@ -1517,10 +1603,10 @@ class MLPProj(torch.nn.Module):
if fl_pos_emb: # NOTE: we only use this for `fl2v` if fl_pos_emb: # NOTE: we only use this for `fl2v`
self.emb_pos = nn.Parameter(torch.zeros(1, 257 * 2, 1280)) 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'): if hasattr(self, 'emb_pos'):
image_embeds = image_embeds + self.emb_pos.to(image_embeds.device) image_embeds = image_embeds + self.emb_pos.to(image_embeds.device)
clip_extra_context_tokens = self.proj(image_embeds) clip_extra_context_tokens = self.proj(image_embeds.to(self.proj[1].weight.dtype)).to(dtype)
return clip_extra_context_tokens return clip_extra_context_tokens
from .s2v.auxi_blocks import MotionEncoder_tc from .s2v.auxi_blocks import MotionEncoder_tc
@@ -1622,47 +1708,22 @@ class AudioInjector_WAN(nn.Module):
class WanModel(torch.nn.Module): class WanModel(torch.nn.Module):
def __init__(self, def __init__(self,
model_type='t2v', model_type='t2v', patch_size=(1, 2, 2),
patch_size=(1, 2, 2), text_len=512, in_dim=16, dim=2048, in_features=5120, out_features=5120,
text_len=512, ffn_dim=8192, ffn2_dim=8192, freq_dim=256, text_dim=4096, out_dim=16,
in_dim=16, num_heads=16, num_layers=32, qk_norm=True, cross_attn_norm=True,
dim=2048, eps=1e-6, attention_mode='sdpa', rope_func='comfy', rms_norm_function='default',
in_features=5120, main_device=torch.device('cuda'), offload_device=torch.device('cpu'),
out_features=5120,
ffn_dim=8192,
ffn2_dim=8192,
freq_dim=256,
text_dim=4096,
out_dim=16,
num_heads=16,
num_layers=32,
qk_norm=True,
cross_attn_norm=True,
eps=1e-6,
attention_mode='sdpa',
rope_func='comfy',
rms_norm_function='default',
main_device=torch.device('cuda'),
offload_device=torch.device('cpu'),
dtype=torch.float16, dtype=torch.float16,
teacache_coefficients=[], teacache_coefficients=[], magcache_ratios=[],
magcache_ratios=[], vace_layers=None, vace_in_dim=None,
vace_layers=None, inject_sample_info=False, add_ref_conv=False,
vace_in_dim=None, in_dim_ref_conv=16, add_control_adapter=False, in_dim_control_adapter=24, use_motion_attn=False,
inject_sample_info=False,
add_ref_conv=False,
in_dim_ref_conv=16,
add_control_adapter=False,
in_dim_control_adapter=24,
use_motion_attn=False,
#s2v #s2v
cond_dim=0, cond_dim=0, audio_dim=1024, num_audio_token=4, enable_adain=False, adain_mode="attn_norm",
audio_dim=1024,
num_audio_token=4,
enable_adain=False,
adain_mode="attn_norm",
audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39], audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
zero_timestep=False, zero_timestep=False,
# humo
humo_audio=False, humo_audio=False,
# WanAnimate # WanAnimate
is_wananimate=False, is_wananimate=False,
@@ -1670,8 +1731,8 @@ class WanModel(torch.nn.Module):
# lynx # lynx
lynx_ip_layers=None, lynx_ip_layers=None,
lynx_ref_layers=None, lynx_ref_layers=None,
# ovi # VAP
is_ovi_audio_model=False, is_VAP = False,
# LongCat # LongCat
is_longcat=False, is_longcat=False,
): ):
@@ -1808,7 +1869,7 @@ class WanModel(torch.nn.Module):
nn.SiLU(), nn.SiLU(),
ConvMLP(dim, dim * 4, kernel_size=7, padding=3), ConvMLP(dim, dim * 4, kernel_size=7, padding=3),
) )
self.original_patch_embedding = self.patch_embedding self.original_patch_embedding = self.patch_embedding
self.expanded_patch_embedding = self.patch_embedding self.expanded_patch_embedding = self.patch_embedding
@@ -1825,6 +1886,13 @@ class WanModel(torch.nn.Module):
adaln_tembed_dim = 512 adaln_tembed_dim = 512
self.time_embedding = TimestepEmbedder(t_embed_dim=adaln_tembed_dim, frequency_embedding_size=freq_dim) 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: 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 self.vace_layers = [i for i in range(0, self.num_layers, 2)] if vace_layers is None else vace_layers
@@ -1865,7 +1933,7 @@ class WanModel(torch.nn.Module):
attention_mode=self.attention_mode, rope_func=self.rope_func, rms_norm_function=rms_norm_function, 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, 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, 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) block_idx=i, is_longcat=is_longcat, mot_ref_block=is_VAP and i in VAP_layers)
for i in range(num_layers) for i in range(num_layers)
]) ])
#MTV Crafter #MTV Crafter
@@ -2139,12 +2207,15 @@ class WanModel(torch.nn.Module):
return x.add(residual_out, alpha=strength) 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): 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 patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0]) t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1]) h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2]) w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if mot_ref:
t_start = -t_len
if steps_t is None: if steps_t is None:
steps_t = t_len steps_t = t_len
if steps_h is None: if steps_h is None:
@@ -2226,6 +2297,7 @@ class WanModel(torch.nn.Module):
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None, x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
flashvsr_LQ_latent=None, flashvsr_strength=1.0, flashvsr_LQ_latent=None, flashvsr_strength=1.0,
num_cond_latents=None, num_cond_latents=None,
x_mot_ref=None, mot_ref_context=None, mot_ref_clip_embeds=None, freqs_mot_ref=None,
): ):
r""" r"""
Forward pass through the diffusion model Forward pass through the diffusion model
@@ -2256,6 +2328,7 @@ class WanModel(torch.nn.Module):
if mtv_motion_tokens is not None: if mtv_motion_tokens is not None:
bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1] bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1]
mtv_motion_tokens = torch.cat([mtv_motion_tokens, self.pad_motion_tokens.to(mtv_motion_tokens).expand(bs, motion_seq_len, -1)], dim=-1) mtv_motion_tokens = torch.cat([mtv_motion_tokens, self.pad_motion_tokens.to(mtv_motion_tokens).expand(bs, motion_seq_len, -1)], dim=-1)
mtv_motion_tokens = mtv_motion_tokens.to(self.base_dtype)
# Fantasy Portrait # Fantasy Portrait
adapter_proj = ip_scale = None adapter_proj = ip_scale = None
@@ -2346,6 +2419,18 @@ class WanModel(torch.nn.Module):
d = self.dim // self.num_heads d = self.dim // self.num_heads
freqs_ovi = rope_params(1024, d - 4 * (d // 6), freqs_scaling=0.19676).to(self.main_device) freqs_ovi = rope_params(1024, d - 4 * (d // 6), freqs_scaling=0.19676).to(self.main_device)
x_ovi = x_ovi.to(self.main_device, self.base_dtype) x_ovi = x_ovi.to(self.main_device, self.base_dtype)
# video-as-prompt motion ref
if x_mot_ref is not None:
x_mot_ref = [self.patch_embedding_mot_ref(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x_mot_ref]
grid_sizes_mot_ref = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x_mot_ref])
x_mot_ref = [u.flatten(2).transpose(1, 2) for u in x_mot_ref]
seq_lens_mot_ref = torch.tensor([u.size(1) for u in x_mot_ref], dtype=torch.int32)
x_mot_ref = torch.cat([torch.cat([u, u.new_zeros(1, seq_lens_mot_ref - u.size(1), u.size(2))], dim=1) for u in x_mot_ref])
x_mot_ref = x_mot_ref.to(self.main_device, self.base_dtype)
num_mot_ref = 1
# WanAnimate # WanAnimate
motion_vec = None motion_vec = None
@@ -2456,13 +2541,16 @@ class WanModel(torch.nn.Module):
s2v_ref_latent.shape[4], s2v_ref_latent.shape[4],
t_start=max(30, F + 9), device=x.device, dtype=x.dtype) t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
freqs = torch.cat([freqs, freqs_ref], dim=1) freqs = torch.cat([freqs, freqs_ref], dim=1)
self.cached_freqs = freqs self.cached_freqs = freqs
self.cached_shape = current_shape self.cached_shape = current_shape
self.cached_cond = has_cond self.cached_cond = has_cond
self.cached_rope_k = self.rope_embedder.k self.cached_rope_k = self.rope_embedder.k
self.cached_ntk_alphas = ntk_alphas self.cached_ntk_alphas = ntk_alphas
if x_mot_ref is not None:
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 # Stand-In RoPE frequencies
if x_ip is not None: if x_ip is not None:
# Generate RoPE frequencies for x_ip # Generate RoPE frequencies for x_ip
@@ -2503,6 +2591,10 @@ class WanModel(torch.nn.Module):
time_embed_dtype = self.base_dtype time_embed_dtype = self.base_dtype
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim
e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim
if x_mot_ref is not None:
t_mod_ref = torch.tensor([1], device=t.device, dtype=t.dtype)
e_mot_ref = self.time_embedding_mot_ref(sinusoidal_embedding_1d(self.freq_dim, t_mod_ref.flatten()).to(time_embed_dtype)) # b, dim
e0_mot_ref = self.time_projection_mot_ref(e_mot_ref).unflatten(1, (6, self.dim)) # b, 6, dim
else: else:
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]: if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
@@ -2567,13 +2659,20 @@ class WanModel(torch.nn.Module):
e = e.to(self.offload_device, non_blocking=self.use_non_blocking) e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
# clip vision embedding # clip vision embedding
clip_embed = None clip_embed = clip_embed_mot_ref = None
if clip_fea is not None and hasattr(self, "img_emb"): if clip_fea is not None and hasattr(self, "img_emb"):
clip_fea = clip_fea.to(self.main_device) clip_fea = clip_fea.to(self.main_device)
if self.offload_img_emb: if self.offload_img_emb:
self.img_emb.to(self.main_device) self.img_emb.to(self.main_device)
clip_embed = self.img_emb(clip_fea) # bs x 257 x dim clip_embed = self.img_emb(clip_fea) # bs x 257 x dim
#context = torch.concat([context_clip, context], dim=1) if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
if mot_ref_clip_embeds is not None:
mot_ref_clip_embeds = mot_ref_clip_embeds.to(self.main_device)
if self.offload_img_emb:
self.img_emb.to(self.main_device)
clip_embed_mot_ref = self.img_emb_mot_ref(mot_ref_clip_embeds) # bs x 257 x dim
if self.offload_img_emb: if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking) self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
@@ -2602,6 +2701,11 @@ class WanModel(torch.nn.Module):
context = self.text_embedding( context = self.text_embedding(
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context]).to(text_embed_dtype)) torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context]).to(text_embed_dtype))
if mot_ref_context is not None:
context_mot_ref = mot_ref_context["prompt_embeds"] if not is_uncond else mot_ref_context["negative_prompt_embeds"]
context_mot_ref = self.text_embedding_mot_ref(
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context_mot_ref]).to(text_embed_dtype))
if self.is_longcat: if self.is_longcat:
context[:, tokens:] = 0 context[:, tokens:] = 0
@@ -2844,6 +2948,13 @@ class WanModel(torch.nn.Module):
kwargs['grid_sizes_ovi'] = grid_sizes_ovi kwargs['grid_sizes_ovi'] = grid_sizes_ovi
kwargs['seq_lens_ovi'] = seq_lens_ovi kwargs['seq_lens_ovi'] = seq_lens_ovi
kwargs['freqs_ovi'] = freqs_ovi kwargs['freqs_ovi'] = freqs_ovi
if x_mot_ref is not None:
kwargs['context_mot_ref'] = context_mot_ref
kwargs['freqs_mot_ref'] = freqs_mot_ref
kwargs['grid_sizes_mot_ref'] = grid_sizes_mot_ref
kwargs['e_mot_ref'] = e0_mot_ref.to(self.base_dtype)
kwargs['num_mot_ref'] = num_mot_ref
kwargs['clip_embed_mot_ref'] = clip_embed_mot_ref
if vace_data is not None: if vace_data is not None:
@@ -2931,7 +3042,9 @@ class WanModel(torch.nn.Module):
if b in self.slg_blocks and is_uncond: if b in self.slg_blocks and is_uncond:
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent: if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
continue continue
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, **kwargs) #run block # ====run block start=====
x, x_ip, lynx_ref_feature, x_ovi, x_mot_ref = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_mot_ref=x_mot_ref, **kwargs)
# ====end run block=====
if self.audio_injector is not None and s2v_audio_input is not None: if self.audio_injector is not None and s2v_audio_input is not None:
x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v
if block.has_face_fuser_block and motion_vec is not None: if block.has_face_fuser_block and motion_vec is not None: