Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7a0da7708e | ||
|
|
a51a53d5b7 | ||
|
|
977f4a5c3a | ||
|
|
2a45675498 | ||
|
|
ea414c54ac | ||
|
|
0013ae0ece | ||
|
|
0e904e6035 |
File diff suppressed because it is too large
Load Diff
@@ -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 = {
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user