Init VAP
This commit is contained in:
@@ -22,6 +22,30 @@ offload_device = mm.unet_offload_device()
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
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:
|
||||
@classmethod
|
||||
@@ -2206,6 +2230,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAnimateEmbeds": WanVideoAnimateEmbeds,
|
||||
"WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents,
|
||||
"WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE,
|
||||
"WanVideoAddVideoPromptEmbeds": WanVideoAddVideoPromptEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
|
||||
@@ -1346,6 +1346,7 @@ class WanVideoModelLoader:
|
||||
"lynx_ip_layers": lynx_ip_layers,
|
||||
"lynx_ref_layers": lynx_ref_layers,
|
||||
"is_longcat": dim == 4096,
|
||||
"is_VAP": True if "patch_embedding_mot_ref.weight" in sd else False
|
||||
|
||||
}
|
||||
|
||||
|
||||
+21
-1
@@ -289,6 +289,7 @@ class WanVideoSampler:
|
||||
phantom_latents = fun_ref_image = ATI_tracks = None
|
||||
add_cond = attn_cond = attn_cond_neg = noise_pred_flipped = None
|
||||
humo_audio = humo_audio_neg = None
|
||||
image_cond_mot_ref = None
|
||||
|
||||
#I2V
|
||||
image_cond = image_embeds.get("image_embeds", None)
|
||||
@@ -476,6 +477,22 @@ class WanVideoSampler:
|
||||
phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0)
|
||||
phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0)
|
||||
|
||||
# Video-as-prompt (VAP)
|
||||
mot_ref_clip_embeds = mot_ref_context = None
|
||||
video_prompt_embeds = image_embeds.get("video_prompt_embeds", None)
|
||||
if video_prompt_embeds is not None:
|
||||
image_cond_mot_ref = video_prompt_embeds.get("image_embeds", None)
|
||||
print("image_cond_mot_ref shape:", image_cond_mot_ref.shape)
|
||||
image_cond_mask_ = video_prompt_embeds.get("mask", None)
|
||||
if image_cond_mask_ is not None:
|
||||
image_cond_mot_ref = torch.cat([image_cond_mask_, image_cond_mot_ref])
|
||||
latents_mot_ref = video_prompt_embeds.get("video_prompt_latents", None)
|
||||
print("latents_mot_ref shape:", latents_mot_ref.shape)
|
||||
x_mot_ref = torch.cat([latents_mot_ref, image_cond_mot_ref], dim=0)
|
||||
mot_ref_context = video_prompt_embeds.get("text_embeds", None)
|
||||
mot_ref_clip_embeds = video_prompt_embeds.get("clip_context", None)
|
||||
print("x_mot_ref shape:", x_mot_ref.shape)
|
||||
|
||||
# CLIP image features
|
||||
clip_fea = image_embeds.get("clip_context", None)
|
||||
if clip_fea is not None:
|
||||
@@ -1395,7 +1412,10 @@ class WanVideoSampler:
|
||||
"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_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.to(z)] if image_cond_mot_ref is not None else None, # 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
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
|
||||
+166
-50
@@ -747,6 +747,34 @@ 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 self.o(x)
|
||||
|
||||
class WanT2VCrossAttentionMOTRef(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
|
||||
self.ip_adapter = None
|
||||
self.k_fusion = None
|
||||
|
||||
def forward(self, x, context, grid_sizes=None, clip_embed=None, rope_func="comfy", **kwargs):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
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):
|
||||
|
||||
@@ -897,7 +925,7 @@ class WanAttentionBlock(nn.Module):
|
||||
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",
|
||||
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__()
|
||||
self.dim = out_features
|
||||
self.ffn_dim = ffn_dim
|
||||
@@ -913,6 +941,7 @@ class WanAttentionBlock(nn.Module):
|
||||
self.dense_block = False
|
||||
self.dense_attention_mode = "sageattn"
|
||||
self.block_idx = block_idx
|
||||
self.mot_ref_block = mot_ref_block
|
||||
|
||||
self.kv_cache = None
|
||||
self.use_motion_attn = use_motion_attn
|
||||
@@ -952,6 +981,16 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
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 = WanT2VCrossAttentionMOTRef(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
|
||||
if use_humo_audio_attn:
|
||||
self.audio_cross_attn_wrapper = AudioCrossAttentionWrapper(in_features, out_features, num_heads, qk_norm, eps, kv_dim=1536)
|
||||
@@ -1045,7 +1084,8 @@ class WanAttentionBlock(nn.Module):
|
||||
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,
|
||||
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"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, C]
|
||||
@@ -1081,6 +1121,17 @@ class WanAttentionBlock(nn.Module):
|
||||
input_x = torch.concat([input_x, input_x_ip], dim=1)
|
||||
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:
|
||||
#import einops
|
||||
# shift_msa_mot_ref, scale_msa_mot_ref, gate_msa_mot_ref, shift_mlp_mot_ref, scale_mlp_mot_ref, gate_mlp_mot_ref = self.get_mod(e_mot_ref.to(x.device), self.modulation_mot_ref)
|
||||
# norm_x_mot_ref = einops.rearrange(self.norm1_mot_ref(x_mot_ref.to(scale_msa_mot_ref.dtype)), 'b (n t) c -> b n t c', n=num_mot_ref)
|
||||
# input_x_mot_ref = self.modulate(norm_x_mot_ref, shift_msa_mot_ref, scale_msa_mot_ref).to(input_dtype)
|
||||
# input_x_mot_ref = einops.rearrange(input_x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref)
|
||||
shift_msa_mot_ref, scale_msa_mot_ref, gate_msa_mot_ref, shift_mlp_mot_ref, scale_mlp_mot_ref, gate_mlp_mot_ref = self.get_mod(e_mot_ref.to(x.device), self.modulation_mot_ref)
|
||||
input_x_mot_ref = self.modulate(self.norm1_mot_ref(x_mot_ref.to(shift_msa_mot_ref.dtype)), shift_msa_mot_ref, scale_msa_mot_ref).to(input_dtype)
|
||||
|
||||
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)
|
||||
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)
|
||||
else:
|
||||
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":
|
||||
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":
|
||||
q, k = apply_rope_comfy_chunked(q, k, freqs)
|
||||
elif self.rope_func == "mocha":
|
||||
@@ -1196,6 +1252,14 @@ class WanAttentionBlock(nn.Module):
|
||||
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
|
||||
# merge x_cond and x_noise
|
||||
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
|
||||
elif use_mot_ref:
|
||||
y_temp = self.self_attn_mot_ref.forward(
|
||||
torch.cat([q, q_mot_ref], dim=1),
|
||||
torch.cat([k, k_mot_ref], dim=1),
|
||||
torch.cat([v, v_mot_ref], dim=1),
|
||||
seq_lens
|
||||
)
|
||||
y, y_mot_ref = torch.split(y_temp, [q.shape[1], q_mot_ref.shape[1]], dim=1)
|
||||
else:
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale)
|
||||
|
||||
@@ -1227,8 +1291,13 @@ class WanAttentionBlock(nn.Module):
|
||||
else:
|
||||
if not is_longcat:
|
||||
x = x.addcmul(y, gate_msa)
|
||||
if use_mot_ref:
|
||||
#x_mot_ref = x_mot_ref.addcmul(einops.rearrange(y_mot_ref, 'b (n t) c -> b n t c', n=num_mot_ref), gate_msa_mot_ref)
|
||||
#x_mot_ref = einops.rearrange(x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref)
|
||||
x_mot_ref = x_mot_ref.addcmul(y_mot_ref, gate_msa_mot_ref)
|
||||
else:
|
||||
x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C)
|
||||
|
||||
del y, gate_msa
|
||||
|
||||
# cross-attention & ffn function
|
||||
@@ -1263,6 +1332,10 @@ class WanAttentionBlock(nn.Module):
|
||||
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)
|
||||
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
|
||||
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,
|
||||
@@ -1270,7 +1343,7 @@ class WanAttentionBlock(nn.Module):
|
||||
x = x.add(x_audio, alpha=audio_scale)
|
||||
|
||||
# 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.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)
|
||||
|
||||
@@ -1317,7 +1390,20 @@ class WanAttentionBlock(nn.Module):
|
||||
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))
|
||||
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 = einops.rearrange(self.norm2_mot_ref(x_mot_ref.to(shift_mlp_mot_ref.dtype)), 'b (n t) c -> b n t c', n=num_mot_ref)
|
||||
# mod_x_mot_ref = torch.addcmul(shift_mlp_mot_ref, norm2_x_mot_ref, 1 + scale_mlp_mot_ref)
|
||||
# mod_x_mot_ref = einops.rearrange(mod_x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref)
|
||||
# x_ffn_mot_ref = self.ffn_mot_ref(mod_x_mot_ref.to(input_dtype))
|
||||
# x_ffn_mot_ref = einops.rearrange(x_ffn_mot_ref, 'b (n t) c -> b n t c', n=num_mot_ref)
|
||||
# x_mot_ref = x_mot_ref.addcmul(x_ffn_mot_ref, gate_mlp_mot_ref)
|
||||
# x_mot_ref = einops.rearrange(x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref)
|
||||
norm2_x_mot_ref = self.norm2_mot_ref(x_mot_ref.to(shift_mlp_mot_ref.dtype))
|
||||
mod_x_mot_ref = torch.addcmul(shift_mlp_mot_ref, norm2_x_mot_ref, 1 + scale_mlp_mot_ref)
|
||||
x_ffn_mot_ref = self.ffn_mot_ref(mod_x_mot_ref.to(input_dtype))
|
||||
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()
|
||||
def split_cross_attn_ffn(self, x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None):
|
||||
@@ -1520,7 +1606,7 @@ class MLPProj(torch.nn.Module):
|
||||
def forward(self, image_embeds):
|
||||
if hasattr(self, 'emb_pos'):
|
||||
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(image_embeds.dtype)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
from .s2v.auxi_blocks import MotionEncoder_tc
|
||||
@@ -1622,47 +1708,22 @@ class AudioInjector_WAN(nn.Module):
|
||||
|
||||
class WanModel(torch.nn.Module):
|
||||
def __init__(self,
|
||||
model_type='t2v',
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=512,
|
||||
in_dim=16,
|
||||
dim=2048,
|
||||
in_features=5120,
|
||||
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'),
|
||||
model_type='t2v', patch_size=(1, 2, 2),
|
||||
text_len=512, in_dim=16, dim=2048, in_features=5120, 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,
|
||||
teacache_coefficients=[],
|
||||
magcache_ratios=[],
|
||||
vace_layers=None,
|
||||
vace_in_dim=None,
|
||||
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,
|
||||
teacache_coefficients=[], magcache_ratios=[],
|
||||
vace_layers=None, vace_in_dim=None,
|
||||
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
|
||||
cond_dim=0,
|
||||
audio_dim=1024,
|
||||
num_audio_token=4,
|
||||
enable_adain=False,
|
||||
adain_mode="attn_norm",
|
||||
cond_dim=0, 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],
|
||||
zero_timestep=False,
|
||||
# humo
|
||||
humo_audio=False,
|
||||
# WanAnimate
|
||||
is_wananimate=False,
|
||||
@@ -1670,8 +1731,8 @@ class WanModel(torch.nn.Module):
|
||||
# lynx
|
||||
lynx_ip_layers=None,
|
||||
lynx_ref_layers=None,
|
||||
# ovi
|
||||
is_ovi_audio_model=False,
|
||||
# VAP
|
||||
is_VAP = False,
|
||||
# LongCat
|
||||
is_longcat=False,
|
||||
):
|
||||
@@ -1808,6 +1869,9 @@ class WanModel(torch.nn.Module):
|
||||
nn.SiLU(),
|
||||
ConvMLP(dim, dim * 4, kernel_size=7, padding=3),
|
||||
)
|
||||
|
||||
if is_VAP:
|
||||
self.patch_embedding_mot_ref = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
|
||||
self.original_patch_embedding = self.patch_embedding
|
||||
self.expanded_patch_embedding = self.patch_embedding
|
||||
@@ -1826,6 +1890,12 @@ class WanModel(torch.nn.Module):
|
||||
self.time_embedding = TimestepEmbedder(t_embed_dim=adaln_tembed_dim, frequency_embedding_size=freq_dim)
|
||||
|
||||
|
||||
if is_VAP:
|
||||
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)
|
||||
|
||||
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_in_dim = self.in_dim if vace_in_dim is None else vace_in_dim
|
||||
@@ -1859,13 +1929,15 @@ class WanModel(torch.nn.Module):
|
||||
else:
|
||||
cross_attn_type = 'no_cross_attn'
|
||||
|
||||
VAP_layers = [0, 4, 8, 12, 16, 20, 24, 28, 32, 36]
|
||||
|
||||
self.blocks = nn.ModuleList([
|
||||
WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads,
|
||||
qk_norm, cross_attn_norm, eps,
|
||||
attention_mode=self.attention_mode, rope_func=self.rope_func, rms_norm_function=rms_norm_function,
|
||||
use_motion_attn=(i % 4 == 0 and use_motion_attn), use_humo_audio_attn=self.humo_audio,
|
||||
face_fuser_block = (i % 5 == 0 and is_wananimate), lynx_ip_layers=lynx_ip_layers, lynx_ref_layers=lynx_ref_layers,
|
||||
block_idx=i, is_longcat=is_longcat)
|
||||
block_idx=i, is_longcat=is_longcat, mot_ref_block=i in VAP_layers and is_VAP)
|
||||
for i in range(num_layers)
|
||||
])
|
||||
#MTV Crafter
|
||||
@@ -2139,12 +2211,15 @@ class WanModel(torch.nn.Module):
|
||||
return x.add(residual_out, alpha=strength)
|
||||
|
||||
|
||||
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
|
||||
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None, mot=False):
|
||||
patch_size = self.patch_size
|
||||
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
|
||||
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
|
||||
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
|
||||
|
||||
if mot:
|
||||
t_start = -t_len
|
||||
|
||||
if steps_t is None:
|
||||
steps_t = t_len
|
||||
if steps_h is None:
|
||||
@@ -2226,6 +2301,7 @@ class WanModel(torch.nn.Module):
|
||||
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
|
||||
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
|
||||
num_cond_latents=None,
|
||||
x_mot_ref=None, mot_ref_context=None, mot_ref_clip_embeds=None,
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -2256,6 +2332,7 @@ class WanModel(torch.nn.Module):
|
||||
if mtv_motion_tokens is not None:
|
||||
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 = mtv_motion_tokens.to(self.base_dtype)
|
||||
|
||||
# Fantasy Portrait
|
||||
adapter_proj = ip_scale = None
|
||||
@@ -2346,6 +2423,18 @@ class WanModel(torch.nn.Module):
|
||||
d = self.dim // self.num_heads
|
||||
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)
|
||||
|
||||
# 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
|
||||
motion_vec = None
|
||||
@@ -2456,13 +2545,16 @@ class WanModel(torch.nn.Module):
|
||||
s2v_ref_latent.shape[4],
|
||||
t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
|
||||
freqs = torch.cat([freqs, freqs_ref], dim=1)
|
||||
|
||||
self.cached_freqs = freqs
|
||||
self.cached_shape = current_shape
|
||||
self.cached_cond = has_cond
|
||||
self.cached_rope_k = self.rope_embedder.k
|
||||
self.cached_ntk_alphas = ntk_alphas
|
||||
|
||||
if x_mot_ref is not None:
|
||||
freqs_mot_ref = self.rope_encode_comfy(F, H, W, mot=True, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype)
|
||||
|
||||
|
||||
# Stand-In RoPE frequencies
|
||||
if x_ip is not None:
|
||||
# Generate RoPE frequencies for x_ip
|
||||
@@ -2503,6 +2595,10 @@ class WanModel(torch.nn.Module):
|
||||
time_embed_dtype = self.base_dtype
|
||||
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
|
||||
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:
|
||||
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
|
||||
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
|
||||
@@ -2592,6 +2688,11 @@ class WanModel(torch.nn.Module):
|
||||
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))
|
||||
|
||||
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:
|
||||
context[:, tokens:] = 0
|
||||
|
||||
@@ -2609,13 +2710,19 @@ class WanModel(torch.nn.Module):
|
||||
else:
|
||||
context = None
|
||||
|
||||
clip_embed = None
|
||||
clip_embed = clip_embed_mot_ref = None
|
||||
if clip_fea is not None and hasattr(self, "img_emb"):
|
||||
clip_fea = clip_fea.to(self.main_device)
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.main_device)
|
||||
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:
|
||||
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
@@ -2841,6 +2948,13 @@ class WanModel(torch.nn.Module):
|
||||
kwargs['grid_sizes_ovi'] = grid_sizes_ovi
|
||||
kwargs['seq_lens_ovi'] = seq_lens_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:
|
||||
@@ -2928,7 +3042,9 @@ class WanModel(torch.nn.Module):
|
||||
if b in self.slg_blocks and is_uncond:
|
||||
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
|
||||
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:
|
||||
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:
|
||||
|
||||
Reference in New Issue
Block a user