@@ -0,0 +1,65 @@
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch
|
||||
import math
|
||||
|
||||
class FeedForwardSwiGLU(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
hidden_dim: int,
|
||||
multiple_of: int = 256,
|
||||
):
|
||||
super().__init__()
|
||||
hidden_dim = int(2 * hidden_dim / 3)
|
||||
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)
|
||||
|
||||
self.dim = dim
|
||||
self.hidden_dim = hidden_dim
|
||||
self.w1 = nn.Linear(dim, hidden_dim, bias=False)
|
||||
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
|
||||
self.w3 = nn.Linear(dim, hidden_dim, bias=False)
|
||||
|
||||
def forward(self, x):
|
||||
return self.w2(F.silu(self.w1(x)) * self.w3(x))
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
def __init__(self, t_embed_dim, frequency_embedding_size=256):
|
||||
super().__init__()
|
||||
self.t_embed_dim = t_embed_dim
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, t_embed_dim, bias=True),
|
||||
nn.SiLU(),
|
||||
nn.Linear(t_embed_dim, t_embed_dim, bias=True),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
:param t: a 1-D Tensor of N indices, one per batch element.
|
||||
These may be fractional.
|
||||
:param dim: the dimension of the output.
|
||||
:param max_period: controls the minimum frequency of the embeddings.
|
||||
:return: an (N, D) Tensor of positional embeddings.
|
||||
"""
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half)
|
||||
freqs = freqs.to(device=t.device)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
def forward(self, t, dtype):
|
||||
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
|
||||
if t_freq.dtype != dtype:
|
||||
t_freq = t_freq.to(dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
File diff suppressed because it is too large
Load Diff
+11
-2
@@ -1174,8 +1174,13 @@ class WanVideoModelLoader:
|
||||
in_features = sd["blocks.0.self_attn.k.weight"].shape[1]
|
||||
out_features = sd["blocks.0.self_attn.k.weight"].shape[0]
|
||||
log.info(f"Detected model in_channels: {in_channels}")
|
||||
ffn_dim = sd["blocks.0.ffn.0.bias"].shape[0]
|
||||
ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1]
|
||||
|
||||
if "blocks.0.ffn.0.bias" in sd:
|
||||
ffn_dim = sd["blocks.0.ffn.0.bias"].shape[0]
|
||||
ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1]
|
||||
else:
|
||||
ffn_dim = sd["blocks.0.ffn.w1.weight"].shape[0]
|
||||
ffn2_dim = sd["blocks.0.ffn.w1.weight"].shape[1]
|
||||
|
||||
patch_size=(1, 2, 2)
|
||||
if "patch_embedding.0.weight" in sd:
|
||||
@@ -1222,6 +1227,9 @@ class WanVideoModelLoader:
|
||||
num_layers = 30
|
||||
out_dim = 48
|
||||
model_type = "t2v" #5B no img crossattn
|
||||
elif dim == 4096: #longcat
|
||||
num_heads = 32
|
||||
num_layers = 48
|
||||
else: #1.3B
|
||||
num_heads = 12
|
||||
num_layers = 30
|
||||
@@ -1335,6 +1343,7 @@ class WanVideoModelLoader:
|
||||
"rms_norm_function": rms_norm_function,
|
||||
"lynx_ip_layers": lynx_ip_layers,
|
||||
"lynx_ref_layers": lynx_ref_layers,
|
||||
"is_longcat": dim == 4096,
|
||||
|
||||
}
|
||||
|
||||
|
||||
+5
-2
@@ -1695,7 +1695,7 @@ class WanVideoSampler:
|
||||
current_step_percentage = idx / len(timesteps)
|
||||
|
||||
timestep = torch.tensor([t]).to(device)
|
||||
if is_pusa or (is_5b and all_indices):
|
||||
if is_pusa or ((is_5b or transformer.is_longcat) and all_indices):
|
||||
orig_timestep = timestep
|
||||
timestep = timestep.unsqueeze(1).repeat(1, latent_video_length)
|
||||
if extra_latents is not None:
|
||||
@@ -2946,7 +2946,10 @@ class WanVideoSampler:
|
||||
if use_tsr:
|
||||
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
|
||||
|
||||
if len(timestep.shape) != 1 and not is_pusa: #5b
|
||||
if transformer.is_longcat:
|
||||
noise_pred = -noise_pred
|
||||
|
||||
if len(timestep.shape) != 1 and not is_pusa: #5b and longcat
|
||||
# all_indices is a list of indices to skip
|
||||
total_indices = list(range(latent.shape[1]))
|
||||
process_indices = [i for i in total_indices if i not in all_indices]
|
||||
|
||||
+217
-116
@@ -345,10 +345,10 @@ class WanRMSNorm(nn.Module):
|
||||
if use_chunked:
|
||||
return self.forward_chunked(x, num_chunks)
|
||||
else:
|
||||
return self._norm(x) * self.weight
|
||||
return self._norm(x.float()).type_as(x) * self.weight
|
||||
|
||||
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):
|
||||
output = torch.empty_like(x)
|
||||
@@ -392,18 +392,29 @@ class WanFusedRMSNorm(nn.RMSNorm):
|
||||
|
||||
return output
|
||||
|
||||
|
||||
import torch.nn.functional as F
|
||||
class WanLayerNorm(nn.LayerNorm):
|
||||
|
||||
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
|
||||
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
|
||||
|
||||
def forward(self, x):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, C]
|
||||
"""
|
||||
return super().forward(x)
|
||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
||||
origin_dtype = inputs.dtype
|
||||
out = F.layer_norm(
|
||||
inputs.float(),
|
||||
self.normalized_shape,
|
||||
None if self.weight is None else self.weight.float(),
|
||||
None if self.bias is None else self.bias.float() ,
|
||||
self.eps
|
||||
).to(origin_dtype)
|
||||
return out
|
||||
|
||||
# def forward(self, x):
|
||||
# r"""
|
||||
# Args:
|
||||
# x(Tensor): Shape [B, L, C]
|
||||
# """
|
||||
# return super().forward(x)
|
||||
|
||||
|
||||
class WanSelfAttention(nn.Module):
|
||||
@@ -416,10 +427,11 @@ class WanSelfAttention(nn.Module):
|
||||
eps=1e-6,
|
||||
attention_mode="sdpa",
|
||||
rms_norm_function="default",
|
||||
kv_dim=None):
|
||||
kv_dim=None,
|
||||
head_norm=False):
|
||||
assert out_features % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = out_features
|
||||
self.dim = min(in_features, out_features)
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = out_features // num_heads
|
||||
self.qk_norm = qk_norm
|
||||
@@ -442,12 +454,14 @@ class WanSelfAttention(nn.Module):
|
||||
self.v = nn.Linear(in_features, out_features)
|
||||
self.o = nn.Linear(in_features, out_features)
|
||||
|
||||
norm_dim = self.head_dim if head_norm else self.dim
|
||||
|
||||
if rms_norm_function=="pytorch":
|
||||
self.norm_q = WanFusedRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanFusedRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_q = WanFusedRMSNorm(norm_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanFusedRMSNorm(norm_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
else:
|
||||
self.norm_q = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_q = WanRMSNorm(norm_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.norm_k = WanRMSNorm(norm_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
def qkv_fn(self, x):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
@@ -456,6 +470,15 @@ class WanSelfAttention(nn.Module):
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
def qkv_fn_longcat(self, x):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
q = self.q(x).view(b, s, n, d)
|
||||
q = self.norm_q(q)
|
||||
k = self.k(x).view(b, s, n, d)
|
||||
k = self.norm_k(k)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
def qkv_fn_ip(self, x):
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
q = self.norm_q(self.q(x) + self.q_loras(x)).view(b, s, n, d)
|
||||
@@ -659,8 +682,8 @@ class LoRALinearLayer(nn.Module):
|
||||
#region crossattn
|
||||
class WanT2VCrossAttention(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"):
|
||||
super().__init__(in_features, out_features, num_heads, qk_norm, eps, kv_dim=kv_dim, rms_norm_function=rms_norm_function)
|
||||
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.attention_mode = attention_mode
|
||||
self.ip_adapter = None
|
||||
self.k_fusion = None
|
||||
@@ -671,12 +694,19 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, **kwargs):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
# compute query
|
||||
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
|
||||
if d == 4096: #longcat
|
||||
q = self.norm_q(self.q(x).view(b, -1, n, d))
|
||||
else:
|
||||
q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d)
|
||||
|
||||
if nag_context is not None and not is_uncond:
|
||||
x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
|
||||
else:
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
if d == 4096:
|
||||
k = self.norm_k(self.k(context).view(b, -1, n, d))
|
||||
else:
|
||||
k = self.norm_k(self.k(context)).view(b, -1, n, d)
|
||||
|
||||
v = self.v(context).view(b, -1, n, d)
|
||||
|
||||
#EchoShot rope
|
||||
@@ -899,12 +929,15 @@ class WanAttentionBlock(nn.Module):
|
||||
face_fuser_block=False,
|
||||
lynx_ip_layers=None,
|
||||
lynx_ref_layers=None,
|
||||
block_idx=0
|
||||
block_idx=0,
|
||||
# long cat
|
||||
is_longcat = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = out_features
|
||||
self.dim = min(out_features, in_features)
|
||||
self.ffn_dim = ffn_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = out_features // num_heads
|
||||
self.qk_norm = qk_norm
|
||||
self.cross_attn_norm = cross_attn_norm
|
||||
self.eps = eps
|
||||
@@ -921,8 +954,9 @@ class WanAttentionBlock(nn.Module):
|
||||
self.has_face_fuser_block = face_fuser_block
|
||||
|
||||
# layers
|
||||
self.norm1 = WanLayerNorm(out_features, eps)
|
||||
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function)
|
||||
self.norm1 = WanLayerNorm(self.dim, eps)
|
||||
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function,
|
||||
head_norm=is_longcat)
|
||||
|
||||
# MTV Crafter motion attn
|
||||
if self.use_motion_attn:
|
||||
@@ -931,14 +965,26 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
if cross_attn_type != "no_cross_attn":
|
||||
self.norm3 = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](in_features, out_features, num_heads, qk_norm, eps, rms_norm_function=rms_norm_function)
|
||||
self.norm2 = WanLayerNorm(out_features, eps)
|
||||
self.ffn = nn.Sequential(
|
||||
nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(ffn2_dim, out_features))
|
||||
self.cross_attn = WAN_CROSSATTENTION_CLASSES[cross_attn_type](in_features, out_features, num_heads, qk_norm, eps, rms_norm_function=rms_norm_function,
|
||||
head_norm=is_longcat)
|
||||
self.norm2 = WanLayerNorm(self.dim, eps)
|
||||
|
||||
if not is_longcat:
|
||||
self.ffn = nn.Sequential(
|
||||
nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(ffn2_dim, out_features))
|
||||
else:
|
||||
from ...LongCat.layers import FeedForwardSwiGLU
|
||||
mlp_ratio = 4
|
||||
self.ffn = FeedForwardSwiGLU(dim=self.dim, hidden_dim=int(self.dim * mlp_ratio))
|
||||
|
||||
# modulation
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
|
||||
if not is_longcat:
|
||||
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
|
||||
else:
|
||||
adaln_tembed_dim = 512
|
||||
self.modulation = nn.Sequential(nn.SiLU(), nn.Linear(adaln_tembed_dim, 6 * self.dim, bias=True))
|
||||
|
||||
self.seg_idx = None
|
||||
|
||||
# HuMo audio cross-attn
|
||||
@@ -962,9 +1008,11 @@ class WanAttentionBlock(nn.Module):
|
||||
if self.block_idx % 2 == 0:
|
||||
self.cross_attn.ip_adapter = WanLynxIPCrossAttention(cross_attention_dim=2048, dim=self.dim, n_registers=0, bias=False)
|
||||
|
||||
#@torch.compiler.disable()
|
||||
def get_mod(self, e, modulation):
|
||||
if e.dim() == 3:
|
||||
if e.shape[-1] == 512:
|
||||
e = self.modulation(e)
|
||||
return e.unsqueeze(2).chunk(6, dim=-1)
|
||||
return (modulation + e).chunk(6, dim=1) # 1, 6, dim
|
||||
elif e.dim() == 4:
|
||||
e_mod = modulation.unsqueeze(2) + e
|
||||
@@ -1029,7 +1077,7 @@ class WanAttentionBlock(nn.Module):
|
||||
mtv_motion_tokens=None, mtv_motion_rotary_emb=None, mtv_strength=1.0, mtv_freqs=None, #mtv crafter
|
||||
humo_audio_input=None, humo_audio_scale=1.0, #humo audio
|
||||
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 #ovi
|
||||
x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None, #ovi
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
@@ -1048,7 +1096,14 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
|
||||
del e
|
||||
input_x = self.modulate(self.norm1(x), shift_msa, scale_msa, seg_idx=self.seg_idx)
|
||||
B, N, C = x.shape
|
||||
T = num_latent_frames
|
||||
is_longcat = C == 4096
|
||||
if is_longcat:
|
||||
input_x = self.modulate(self.norm1(x.view(B, T, -1, C).float()).to(x.dtype), shift_msa, scale_msa, seg_idx=self.seg_idx).view(B, N, C)
|
||||
else:
|
||||
input_x = self.modulate(self.norm1(x), shift_msa, scale_msa, seg_idx=self.seg_idx)
|
||||
|
||||
del shift_msa, scale_msa
|
||||
|
||||
if x_ip is not None:
|
||||
@@ -1098,7 +1153,10 @@ class WanAttentionBlock(nn.Module):
|
||||
elif self.rope_func == "comfy_chunked":
|
||||
q_ip, k_ip = apply_rope_comfy_chunked(q_ip, k_ip, freqs_ip)
|
||||
else:
|
||||
q, k, v = self.self_attn.qkv_fn(input_x)
|
||||
if is_longcat:
|
||||
q, k, v = self.self_attn.qkv_fn_longcat(input_x)
|
||||
else:
|
||||
q, k, v = self.self_attn.qkv_fn(input_x)
|
||||
if self.rope_func == "comfy":
|
||||
q, k = apply_rope_comfy(q, k, freqs)
|
||||
elif self.rope_func == "comfy_chunked":
|
||||
@@ -1193,7 +1251,10 @@ class WanAttentionBlock(nn.Module):
|
||||
y = torch.cat(z, dim=1)
|
||||
x = x.add(y)
|
||||
else:
|
||||
x = x.addcmul(y, gate_msa)
|
||||
if not is_longcat:
|
||||
x = x.addcmul(y, gate_msa)
|
||||
else:
|
||||
x = x + (y.view(B, -1, N//T, C).float() * gate_msa).view(B, -1, C).to(x.dtype)
|
||||
del y, gate_msa
|
||||
|
||||
# cross-attention & ffn function
|
||||
@@ -1210,8 +1271,6 @@ class WanAttentionBlock(nn.Module):
|
||||
y = self.audio_block.ffn(torch.addcmul(shift_mlp_ovi, self.audio_block.norm2(x_ovi), 1 + scale_mlp_ovi))
|
||||
x_ovi = x_ovi.addcmul(y, gate_mlp_ovi)
|
||||
|
||||
assert not torch.equal(og_ovi_x, x_ovi), "Audio should be changed after cross-attention!"
|
||||
|
||||
# video
|
||||
x = x + self.cross_attn(self.norm3(x), context, grid_sizes,
|
||||
src_freqs=freqs,
|
||||
@@ -1226,20 +1285,58 @@ class WanAttentionBlock(nn.Module):
|
||||
raise NotImplementedError("nag_context is not supported in split_cross_attn_ffn")
|
||||
x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes)
|
||||
else:
|
||||
x = self.cross_attn_ffn(x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
|
||||
audio_proj, audio_scale, num_latent_frames, nag_params, nag_context, is_uncond,
|
||||
multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs,
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale, zero_timestep=zero_timestep,
|
||||
mtv_freqs=mtv_freqs, mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength,
|
||||
humo_audio_input=humo_audio_input, humo_audio_scale=humo_audio_scale, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, original_seq_len=original_seq_len
|
||||
)
|
||||
else:
|
||||
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||
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)
|
||||
# MultiTalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
|
||||
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
|
||||
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:
|
||||
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
|
||||
x = x.add(x_motion, alpha=mtv_strength)
|
||||
|
||||
# HuMo Audio Cross-Attention
|
||||
if humo_audio_input is not None:
|
||||
x = self.audio_cross_attn_wrapper(x, humo_audio_input, grid_sizes, humo_audio_scale)
|
||||
|
||||
# ffn
|
||||
if self.rope_func == "comfy_chunked":
|
||||
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
x_ffn = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
else:
|
||||
y = self.ffn(torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp))
|
||||
x = x.addcmul(y, gate_mlp)
|
||||
del gate_mlp
|
||||
if zero_timestep:
|
||||
norm2_x = self.norm2(x)
|
||||
parts = []
|
||||
for i in range(2):
|
||||
parts.append(norm2_x[:, self.seg_idx[i]:self.seg_idx[i + 1]] *
|
||||
(1 + scale_mlp[:, i:i + 1]) + shift_mlp[:, i:i + 1])
|
||||
norm2_x = torch.cat(parts, dim=1)
|
||||
x_ffn = self.ffn(norm2_x)
|
||||
else:
|
||||
if not is_longcat:
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp)
|
||||
else:
|
||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()).to(x.dtype), 1 + scale_mlp).view(B, -1, C)
|
||||
x_ffn = self.ffn(mod_x)
|
||||
del shift_mlp, scale_mlp
|
||||
|
||||
# gate_mlp
|
||||
if zero_timestep:
|
||||
z = []
|
||||
for i in range(2):
|
||||
z.append(x_ffn[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_mlp[:, i:i + 1])
|
||||
x_ffn = torch.cat(z, dim=1)
|
||||
x = x.add(x_ffn)
|
||||
else:
|
||||
if not is_longcat:
|
||||
x = x.addcmul(x_ffn, gate_mlp)
|
||||
else:
|
||||
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).view(B, -1, C).to(x.dtype)
|
||||
del gate_mlp
|
||||
|
||||
if x_ip is not None: #stand-in
|
||||
x_ip = x_ip.addcmul(y_ip, gate_msa_ip)
|
||||
@@ -1248,58 +1345,6 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
return x, x_ip, lynx_ref_feature, x_ovi
|
||||
|
||||
|
||||
def cross_attn_ffn(self, x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
|
||||
audio_proj, audio_scale, num_latent_frames, nag_params,
|
||||
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num,
|
||||
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale, zero_timestep, mtv_freqs, mtv_motion_tokens, mtv_motion_rotary_emb, mtv_strength,
|
||||
humo_audio_input, humo_audio_scale, lynx_x_ip, lynx_ip_scale, original_seq_len):
|
||||
|
||||
x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed,
|
||||
audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||
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, )
|
||||
# MultiTalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
|
||||
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
|
||||
x = x + x_audio * 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:
|
||||
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
|
||||
x = x + x_motion * mtv_strength
|
||||
|
||||
# HuMo Audio Cross-Attention
|
||||
if humo_audio_input is not None:
|
||||
x = self.audio_cross_attn_wrapper(x, humo_audio_input, grid_sizes, humo_audio_scale)
|
||||
|
||||
if self.rope_func == "comfy_chunked" and not zero_timestep:
|
||||
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
|
||||
else:
|
||||
norm2_x = self.norm2(x)
|
||||
if zero_timestep:
|
||||
parts = []
|
||||
for i in range(2):
|
||||
parts.append(norm2_x[:, self.seg_idx[i]:self.seg_idx[i + 1]] *
|
||||
(1 + scale_mlp[:, i:i + 1]) + shift_mlp[:, i:i + 1])
|
||||
norm2_x = torch.cat(parts, dim=1)
|
||||
y = self.ffn(norm2_x)
|
||||
else:
|
||||
input_x = torch.addcmul(shift_mlp, norm2_x, 1 + scale_mlp)
|
||||
del shift_mlp, scale_mlp, norm2_x
|
||||
y = self.ffn(input_x)
|
||||
if zero_timestep:
|
||||
z = []
|
||||
for i in range(2):
|
||||
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_mlp[:, i:i + 1])
|
||||
y = torch.cat(z, dim=1)
|
||||
x = x.add(y)
|
||||
else:
|
||||
x = x.addcmul(y, gate_mlp)
|
||||
return x
|
||||
|
||||
@torch.compiler.disable()
|
||||
def split_cross_attn_ffn(self, x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None):
|
||||
# Get number of prompts
|
||||
@@ -1443,7 +1488,7 @@ class Head(nn.Module):
|
||||
e = (self.modulation.unsqueeze(2) + e.unsqueeze(1)).chunk(2, dim=1)
|
||||
return [ei.squeeze(1) for ei in e]
|
||||
|
||||
def forward(self, x, e):
|
||||
def forward(self, x, e, **kwargs):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
@@ -1454,6 +1499,37 @@ class Head(nn.Module):
|
||||
x = self.head(self.norm(x).mul_(1 + e[1]).add_(e[0]))
|
||||
return x
|
||||
|
||||
class Head_adaLN(nn.Module):
|
||||
|
||||
def __init__(self, dim, out_dim, patch_size, eps=1e-6, adaln_tembed_dim=512):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.out_dim = out_dim
|
||||
self.patch_size = patch_size
|
||||
self.eps = eps
|
||||
self.adaln_tembed_dim = adaln_tembed_dim
|
||||
|
||||
# layers
|
||||
out_dim = math.prod(patch_size) * out_dim
|
||||
self.norm = WanLayerNorm(dim, eps)
|
||||
self.head = nn.Linear(dim, out_dim)
|
||||
|
||||
# modulation
|
||||
self.modulation = nn.Sequential(nn.SiLU(), nn.Linear(adaln_tembed_dim, 2 * self.dim, bias=True))
|
||||
|
||||
def forward(self, x, e, temp_length):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
e(Tensor): Shape [B, C]
|
||||
"""
|
||||
B, N, C = x.shape
|
||||
T = temp_length
|
||||
self.modulation.to(torch.float32)
|
||||
shift, scale = self.modulation(e).unsqueeze(2).chunk(2, dim=-1) # [B, T, 1, C]
|
||||
return self.head(self.norm(x.view(B, T, -1, C).float()).mul_(1 + scale).add_(shift).view(B, N, C).to(x.dtype))
|
||||
|
||||
|
||||
|
||||
class MLPProj(torch.nn.Module):
|
||||
|
||||
@@ -1622,6 +1698,8 @@ class WanModel(torch.nn.Module):
|
||||
lynx_ref_layers=None,
|
||||
# ovi
|
||||
is_ovi_audio_model=False,
|
||||
# LongCat
|
||||
is_longcat=False,
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
@@ -1744,6 +1822,8 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
self.audio_model = None
|
||||
|
||||
self.is_longcat = is_longcat
|
||||
|
||||
# embeddings
|
||||
if not self.is_ovi_audio_model:
|
||||
self.patch_embedding = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
@@ -1763,9 +1843,14 @@ class WanModel(torch.nn.Module):
|
||||
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
|
||||
nn.Linear(dim, dim))
|
||||
|
||||
self.time_embedding = nn.Sequential(
|
||||
nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
|
||||
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
|
||||
if not is_longcat:
|
||||
self.time_embedding = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
|
||||
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
|
||||
else:
|
||||
from ...LongCat.layers import TimestepEmbedder
|
||||
adaln_tembed_dim = 512
|
||||
self.time_embedding = TimestepEmbedder(t_embed_dim=adaln_tembed_dim, frequency_embedding_size=freq_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
|
||||
@@ -1805,7 +1890,8 @@ class WanModel(torch.nn.Module):
|
||||
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)
|
||||
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)
|
||||
for i in range(num_layers)
|
||||
])
|
||||
#MTV Crafter
|
||||
@@ -1813,8 +1899,10 @@ class WanModel(torch.nn.Module):
|
||||
self.pad_motion_tokens = torch.zeros(1, 1, 2048)
|
||||
|
||||
# head
|
||||
self.head = Head(dim, out_dim, patch_size, eps)
|
||||
|
||||
if not is_longcat:
|
||||
self.head = Head(dim, out_dim, patch_size, eps)
|
||||
else:
|
||||
self.head = Head_adaLN(dim, out_dim, patch_size, eps, adaln_tembed_dim=512)
|
||||
|
||||
d = self.dim // self.num_heads
|
||||
self.rope_embedder = EmbedND_RifleX(
|
||||
@@ -2372,7 +2460,7 @@ class WanModel(torch.nn.Module):
|
||||
x_ip = ip_image_patch.flatten(2).transpose(1, 2) # [B, N, D]
|
||||
freq_offset = standin_input["freq_offset"]
|
||||
|
||||
if freqs is None: #comfy rope
|
||||
if freqs is None and "comfy" in self.rope_func: #comfy rope
|
||||
current_shape = (F, H, W)
|
||||
|
||||
has_cond = attn_cond is not None
|
||||
@@ -2425,7 +2513,7 @@ class WanModel(torch.nn.Module):
|
||||
freqs = torch.cat([freqs, freqs_motion], dim=1)
|
||||
|
||||
# time embeddings
|
||||
if t.dim() == 2:
|
||||
if t.dim() == 2 and not self.is_longcat:
|
||||
b, f = t.shape
|
||||
expanded_timesteps = True
|
||||
else:
|
||||
@@ -2434,12 +2522,21 @@ class WanModel(torch.nn.Module):
|
||||
if self.zero_timestep:
|
||||
t = torch.cat([t, torch.zeros([1], dtype=t.dtype, device=t.device)])
|
||||
|
||||
time_embed_dtype = self.time_embedding[0].weight.dtype
|
||||
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
|
||||
time_embed_dtype = self.base_dtype
|
||||
if hasattr(self, "time_projection"):
|
||||
time_embed_dtype = self.time_embedding[0].weight.dtype
|
||||
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
|
||||
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
|
||||
else:
|
||||
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
|
||||
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
|
||||
time_embed_dtype = self.base_dtype
|
||||
if len(t.shape) == 1:
|
||||
t = t.unsqueeze(1).expand(-1, F) # [B, T]
|
||||
self.time_embedding.to(torch.float32)
|
||||
e = e0 = self.time_embedding(t.float().flatten(), dtype=torch.float32).reshape(1, F, -1)
|
||||
|
||||
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 self.audio_model is not None:
|
||||
#if t.dim() == 1:
|
||||
@@ -2516,8 +2613,12 @@ class WanModel(torch.nn.Module):
|
||||
context_ovi = self.audio_model.text_embedding(
|
||||
torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context_ovi]).to(text_embed_dtype))
|
||||
|
||||
tokens = context[0].shape[0]
|
||||
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 self.is_longcat:
|
||||
context[:, tokens:] = 0
|
||||
|
||||
# NAG
|
||||
if nag_context is not None:
|
||||
@@ -2928,7 +3029,7 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
x = x[:, :self.original_seq_len]
|
||||
|
||||
x = self.head(x, e.to(x.device))
|
||||
x = self.head(x, e.to(x.device), temp_length=F)
|
||||
|
||||
if x_ovi is not None:
|
||||
x_ovi = self.audio_model.head(x_ovi, e_ovi.to(x_ovi.device))
|
||||
|
||||
@@ -20,6 +20,7 @@ scheduler_list = [
|
||||
"dpm++", "dpm++/beta",
|
||||
"dpm++_sde", "dpm++_sde/beta",
|
||||
"euler", "euler/beta",
|
||||
"longcat_distill_euler",
|
||||
"deis",
|
||||
"lcm", "lcm/beta",
|
||||
"res_multistep",
|
||||
@@ -42,12 +43,24 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
|
||||
elif scheduler in ['euler/beta', 'euler']:
|
||||
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
|
||||
if flowedit_args: #seems to work better
|
||||
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift))
|
||||
elif scheduler in ['euler/beta', 'euler', 'longcat_distill_euler']:
|
||||
if 'longcat' in scheduler:
|
||||
num_distill_sample_steps = 50
|
||||
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, time_shift_type="linear")
|
||||
distill_indices = torch.arange(1, num_distill_sample_steps + 1, dtype=torch.float32)
|
||||
distill_indices = (distill_indices * (1000 // num_distill_sample_steps)).round().long()
|
||||
|
||||
inference_indices = torch.linspace(0, num_distill_sample_steps, steps+1)[:-1]
|
||||
inference_indices = torch.floor(inference_indices).to(torch.int64)
|
||||
|
||||
sigmas = torch.flip(distill_indices, [0])[inference_indices].float() / 1000
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas)
|
||||
else:
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None)
|
||||
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
|
||||
if flowedit_args: #seems to work better
|
||||
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift))
|
||||
else:
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas)
|
||||
elif 'dpm' in scheduler:
|
||||
if 'sde' in scheduler:
|
||||
algorithm_type = "sde-dpmsolver++"
|
||||
|
||||
Reference in New Issue
Block a user