ref latent
This commit is contained in:
@@ -2213,11 +2213,14 @@ class WanVideoSampler:
|
||||
mtv_freqs = mtv_freqs.to(device, dtype)
|
||||
|
||||
#region S2V
|
||||
s2v_audio_input = None
|
||||
s2v_audio_input = s2v_ref_latent = None
|
||||
s2v_audio_embeds = image_embeds.get("audio_embeds", None)
|
||||
if s2v_audio_embeds is not None:
|
||||
log.info(f"Using S2V audio embeddings")
|
||||
s2v_audio_input = s2v_audio_embeds["audio_embed_bucket"].to(device, dtype)
|
||||
s2v_ref_latent = s2v_audio_embeds["ref_latent"]
|
||||
if s2v_ref_latent is not None:
|
||||
s2v_ref_latent = s2v_ref_latent.to(device, dtype)
|
||||
#s2v_audio_input_all_layers = s2v_audio_embeds["audio_encoder_output"]["encoded_audio_all_layers"]
|
||||
print(s2v_audio_input.shape)
|
||||
##print(s2v_audio_input_all_layers[0].shape)
|
||||
@@ -2679,7 +2682,8 @@ class WanVideoSampler:
|
||||
"mtv_motion_rotary_emb": mtv_motion_rotary_emb if mtv_input is not None else None, # MTV-Crafter RoPE
|
||||
"mtv_strength": mtv_strength[idx] if mtv_input is not None else 1.0, # MTV-Crafter scaling
|
||||
"mtv_freqs": mtv_freqs if mtv_input is not None else None, # MTV-Crafter extra RoPE freqs
|
||||
"s2v_audio_input": s2v_audio_input #official speech-to-video
|
||||
"s2v_audio_input": s2v_audio_input, # official speech-to-video audio input
|
||||
"s2v_ref_latent": s2v_ref_latent # official speech-to-video reference latent
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
|
||||
@@ -1190,7 +1190,8 @@ class WanVideoModelLoader:
|
||||
"add_control_adapter": True if "control_adapter.conv.weight" in sd else False,
|
||||
"use_motion_attn": True if "blocks.0.motion_attn.k.weight" in sd else False,
|
||||
"enable_adain": True if "audio_injector.injector_adain_layers.0.linear.weight" in sd else False,
|
||||
"cond_dim": sd["cond_encoder.weight"].shape[1] if "cond_encoder.weight" in sd else 0
|
||||
"cond_dim": sd["cond_encoder.weight"].shape[1] if "cond_encoder.weight" in sd else 0,
|
||||
"zero_timestep": model_type == "s2v",
|
||||
|
||||
}
|
||||
|
||||
|
||||
+12
-11
@@ -52,11 +52,12 @@ class WanVideoAddAudioEmbeds:
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",),
|
||||
"input_fps": ("FLOAT", {"default": 50.0, "min": 1.0, "max": 120.0, "step": 1.0, "tooltip": "Frames per second for the audio"}),
|
||||
"output_fps": ("FLOAT", {"default": 30.0, "min": 1.0, "max": 120.0, "step": 1.0, "tooltip": "Frames per second for the video"}),
|
||||
"bucket_fps": ("FLOAT", {"default": 16.0, "min": 1.0, "max": 120.0, "step": 1.0, "tooltip": "Frames per second for the generated video"}),
|
||||
"frames": ("INT", {"default": 80, "min": 1, "max": 120, "step": 1, "tooltip": "Number of frames to process"})
|
||||
"frames": ("INT", {"default": 81, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames to process"}),
|
||||
},
|
||||
"optional": {
|
||||
"ref_latent": ("LATENT",)
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
@@ -64,15 +65,14 @@ class WanVideoAddAudioEmbeds:
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, input_fps, output_fps, bucket_fps, frames, audio_encoder_output):
|
||||
# Prepare the new audio entry
|
||||
|
||||
#audio_feat = audio_encoder_output["encoded_audio"]
|
||||
#print("audio_feat", audio_feat.shape)
|
||||
def add(self, embeds, frames, audio_encoder_output, ref_latent=None):
|
||||
all_layers = audio_encoder_output["encoded_audio_all_layers"]
|
||||
audio_feat = torch.stack(all_layers, dim=0).squeeze(1) # shape: [num_layers, T, 512]
|
||||
|
||||
print("audio_feat", audio_feat.shape)
|
||||
input_fps = 50
|
||||
output_fps = 30
|
||||
bucket_fps = 16
|
||||
|
||||
if input_fps != output_fps:
|
||||
audio_feat = linear_interpolation(audio_feat, input_fps=input_fps, output_fps=output_fps)
|
||||
@@ -82,7 +82,7 @@ class WanVideoAddAudioEmbeds:
|
||||
audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps(
|
||||
audio_feat,
|
||||
fps=bucket_fps,
|
||||
batch_frames=frames
|
||||
batch_frames=frames-1
|
||||
)
|
||||
|
||||
audio_embed_bucket = audio_embed_bucket.unsqueeze(0)
|
||||
@@ -97,7 +97,8 @@ class WanVideoAddAudioEmbeds:
|
||||
|
||||
new_entry = {
|
||||
"audio_embed_bucket": audio_embed_bucket,
|
||||
"num_repeat": num_repeat
|
||||
"num_repeat": num_repeat,
|
||||
"ref_latent": ref_latent["samples"] if ref_latent is not None else None
|
||||
}
|
||||
updated = dict(embeds)
|
||||
updated["audio_embeds"] = new_entry
|
||||
|
||||
+100
-31
@@ -19,7 +19,7 @@ except:
|
||||
|
||||
from .attention import attention
|
||||
import numpy as np
|
||||
|
||||
from copy import deepcopy
|
||||
from tqdm import tqdm
|
||||
import gc
|
||||
|
||||
@@ -796,8 +796,24 @@ class WanAttentionBlock(nn.Module):
|
||||
e = (self.modulation.unsqueeze(2) + e).chunk(6, dim=1) # 1, 6, 1, dim
|
||||
return [ei.squeeze(1) for ei in e]
|
||||
|
||||
def modulate(self, x, shift_msa, scale_msa):
|
||||
return torch.addcmul(shift_msa, x, 1 + scale_msa)
|
||||
def modulate(self, x, shift_msa, scale_msa, seg_idx=None):
|
||||
"""
|
||||
Modulate x with shift and scale. If seg_idx is provided, apply segmented modulation.
|
||||
"""
|
||||
norm_x = self.norm1(x)
|
||||
if seg_idx is not None:
|
||||
parts = []
|
||||
for i in range(2):
|
||||
part = torch.addcmul(
|
||||
shift_msa[:, i:i + 1],
|
||||
norm_x[:, seg_idx[i]:seg_idx[i + 1]],
|
||||
1 + scale_msa[:, i:i + 1]
|
||||
)
|
||||
parts.append(part)
|
||||
norm_x = torch.cat(parts, dim=1)
|
||||
return norm_x
|
||||
else:
|
||||
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
|
||||
|
||||
def ffn_chunked(self, x, shift_mlp, scale_mlp, num_chunks=4):
|
||||
modulated_input = torch.addcmul(shift_mlp, self.norm2(x), 1 + scale_mlp)
|
||||
@@ -864,9 +880,15 @@ class WanAttentionBlock(nn.Module):
|
||||
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
|
||||
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
|
||||
"""
|
||||
#e = (self.modulation.to(e.device) + e).chunk(6, dim=1)
|
||||
self.zero_timestep = len(e) == 2
|
||||
if self.zero_timestep: #s2v zero timestep
|
||||
self.seg_idx = e[1]
|
||||
self.seg_idx = min(max(0, self.seg_idx), x.size(1))
|
||||
self.seg_idx = [0, self.seg_idx, x.size(1)]
|
||||
e = e[0]
|
||||
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device))
|
||||
input_x = self.modulate(self.norm1(x), shift_msa, scale_msa)
|
||||
input_x = self.modulate(x, shift_msa, scale_msa, seg_idx=self.seg_idx)
|
||||
|
||||
if x_ip is not None:
|
||||
shift_msa_ip, scale_msa_ip, gate_msa_ip, shift_mlp_ip, scale_mlp_ip, gate_mlp_ip = self.get_mod(e_ip.to(x.device))
|
||||
@@ -984,7 +1006,14 @@ class WanAttentionBlock(nn.Module):
|
||||
y[:, -self.cond_size :],
|
||||
)
|
||||
|
||||
x = x.addcmul(y, gate_msa)
|
||||
if self.zero_timestep:
|
||||
z = []
|
||||
for i in range(2):
|
||||
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_msa[:, i:i + 1])
|
||||
y = torch.cat(z, dim=1)
|
||||
x = x.add(y)
|
||||
else:
|
||||
x = x.addcmul(y, gate_msa)
|
||||
|
||||
# cross-attention & ffn function
|
||||
if context is not None:
|
||||
@@ -1037,8 +1066,24 @@ class WanAttentionBlock(nn.Module):
|
||||
if self.rope_func == "comfy_chunked":
|
||||
y = 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)
|
||||
norm2_x = self.norm2(x)
|
||||
if self.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:
|
||||
y = self.ffn(torch.addcmul(shift_mlp, norm2_x, 1 + scale_mlp))
|
||||
if self.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()
|
||||
@@ -1338,6 +1383,7 @@ class WanModel(torch.nn.Module):
|
||||
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
|
||||
):
|
||||
r"""
|
||||
Initialize the diffusion model backbone.
|
||||
@@ -1571,6 +1617,7 @@ class WanModel(torch.nn.Module):
|
||||
need_adain_ont=adain_mode != "attn_norm",
|
||||
)
|
||||
self.adain_mode = adain_mode
|
||||
self.zero_timestep = zero_timestep
|
||||
|
||||
self.trainable_cond_mask = nn.Embedding(3, self.dim)
|
||||
|
||||
@@ -1809,8 +1856,9 @@ class WanModel(torch.nn.Module):
|
||||
mtv_motion_rotary_emb=None,
|
||||
mtv_freqs=None,
|
||||
mtv_strength=1.0,
|
||||
s2v_audio_input=None
|
||||
|
||||
s2v_audio_input=None,
|
||||
s2v_ref_latent=None
|
||||
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -1917,13 +1965,16 @@ class WanModel(torch.nn.Module):
|
||||
fun_camera = self.control_adapter(fun_camera)
|
||||
x = [u + v for u, v in zip(x, fun_camera)]
|
||||
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x])
|
||||
|
||||
grid_sizes = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x])
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.float32)
|
||||
assert seq_lens.max() <= seq_len
|
||||
|
||||
x_len = x[0].shape[1]
|
||||
|
||||
self.original_seq_len = x[0].size(1)
|
||||
|
||||
if add_cond is not None:
|
||||
add_cond = self.add_conv_in(add_cond.to(self.add_conv_in.weight.dtype)).to(x[0].dtype)
|
||||
add_cond = add_cond.flatten(2).transpose(1, 2)
|
||||
@@ -1943,24 +1994,26 @@ class WanModel(torch.nn.Module):
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
seq_len += fun_ref.size(1)
|
||||
F += 1
|
||||
x = [torch.concat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)]
|
||||
x = [torch.cat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)]
|
||||
|
||||
if phantom_ref is not None:
|
||||
phantom_ref_frames = phantom_ref.size(1)
|
||||
phantom_ref = self.original_patch_embedding(phantom_ref.unsqueeze(0).to(torch.float32)).flatten(2).transpose(1, 2).to(x[0].dtype)
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
phantom_ref_seq_len = phantom_ref.size(1)
|
||||
seq_len += phantom_ref_seq_len
|
||||
F += phantom_ref_frames
|
||||
x = [torch.concat([u, phantom_ref.unsqueeze(0)], dim=1) for phantom_ref, u in zip(phantom_ref, x)]
|
||||
end_ref_latent=None
|
||||
if s2v_ref_latent is not None:
|
||||
end_ref_latent = s2v_ref_latent.squeeze(0)
|
||||
elif phantom_ref is not None:
|
||||
end_ref_latent = phantom_ref
|
||||
if end_ref_latent is not None:
|
||||
end_ref_latent_frames = end_ref_latent.size(1)
|
||||
end_ref_latent = self.original_patch_embedding(end_ref_latent.unsqueeze(0).to(torch.float32)).flatten(2).transpose(1, 2).to(x[0].dtype)
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + end_ref_latent_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
end_ref_latent_seq_len = end_ref_latent.size(1)
|
||||
seq_len += end_ref_latent_seq_len
|
||||
F += end_ref_latent_frames
|
||||
x = [torch.cat([u, end_ref_latent.unsqueeze(0)], dim=1) for end_ref_latent, u in zip(end_ref_latent, x)]
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.float32)
|
||||
self.original_seq_len = x[0].size(1)
|
||||
|
||||
assert seq_lens.max() <= seq_len
|
||||
grid_sizes = grid_sizes
|
||||
x = torch.cat([
|
||||
torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))],
|
||||
dim=1) for u in x
|
||||
dim=1) for u in x
|
||||
])
|
||||
|
||||
# StandIn LoRA input
|
||||
@@ -2053,9 +2106,24 @@ class WanModel(torch.nn.Module):
|
||||
else:
|
||||
expanded_timesteps = False
|
||||
|
||||
if self.zero_timestep:
|
||||
t = torch.cat([t, torch.zeros([1], dtype=t.dtype, device=t.device)])
|
||||
|
||||
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(x.dtype)) # b, dim
|
||||
e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim
|
||||
|
||||
#S2V zero timestep
|
||||
if self.zero_timestep:
|
||||
e = e[:-1]
|
||||
zero_e0 = e0[-1:]
|
||||
e0 = e0[:-1]
|
||||
e0 = torch.cat([
|
||||
e0.unsqueeze(2),
|
||||
zero_e0.unsqueeze(2).repeat(e0.size(0), 1, 1, 1)
|
||||
],
|
||||
dim=2)
|
||||
e0 = [e0, self.original_seq_len]
|
||||
|
||||
if x_ip is not None:
|
||||
timestep_ip = torch.zeros_like(t) # [B] with 0s
|
||||
t_ip = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, timestep_ip.flatten()).to(x.dtype)) # b, dim )
|
||||
@@ -2438,15 +2506,16 @@ class WanModel(torch.nn.Module):
|
||||
x = x[:, fun_ref_length:]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if phantom_ref is not None:
|
||||
phantom_ref_length = phantom_ref.size(1)
|
||||
x = x[:, :-phantom_ref_length]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - phantom_ref_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
if end_ref_latent is not None:
|
||||
end_ref_latent_length = end_ref_latent.size(1)
|
||||
x = x[:, :-end_ref_latent_length]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - end_ref_latent_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if attn_cond is not None:
|
||||
x = x[:, :x_len]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
x = x[:, :self.original_seq_len]
|
||||
x = self.head(x, e.to(x.device))
|
||||
x = self.unpatchify(x, grid_sizes) # type: ignore[arg-type]
|
||||
x = [u.float() for u in x]
|
||||
|
||||
Reference in New Issue
Block a user