ref latent

This commit is contained in:
kijai
2025-08-26 21:30:10 +03:00
parent 63d4b6aada
commit 3c79851230
4 changed files with 120 additions and 45 deletions
+6 -2
View File
@@ -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
+2 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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]