Merge branch 's2v'

This commit is contained in:
kijai
2025-08-31 18:05:30 +03:00
13 changed files with 7035 additions and 171 deletions
+3
View File
@@ -12,6 +12,7 @@ from .nodes_model_loading import NODE_CLASS_MAPPINGS as MODEL_LOADING_NODE_CLASS
from .nodes_utility import NODE_CLASS_MAPPINGS as UTILITY_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UTILITY_NODE_DISPLAY_NAME_MAPPINGS
from .cache_methods.nodes_cache import NODE_CLASS_MAPPINGS as NODE_CACHE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as NODE_CACHE_DISPLAY_NAME_MAPPINGS
from .nodes_deprecated import NODE_CLASS_MAPPINGS as DEPRECATED_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS
from .s2v.nodes import NODE_CLASS_MAPPINGS as S2V_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as S2V_NODE_DISPLAY_NAME_MAPPINGS
try:
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
@@ -58,6 +59,7 @@ NODE_CLASS_MAPPINGS.update(NODE_CACHE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(DEPRECATED_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(QWEN_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MTV_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(S2V_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
@@ -75,5 +77,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(NODE_CACHE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(QWEN_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MTV_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(S2V_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+8
View File
@@ -261,6 +261,14 @@ class MultiTalkWav2VecEmbeds:
offset += length
multitalk_audio_features = full_list
# if audio_encoder_output is not None:
# all_layers = audio_encoder_output["encoded_audio_all_layers"]
# audio_feat = torch.stack(all_layers, dim=0).squeeze(1)[1:] # shape: [num_layers, T, 512]
# audio_feat = audio_feat.movedim(0, 1)
# print("audio_feat mean", audio_feat.mean())
# print("audio_feat min max", audio_feat.min(), audio_feat.max())
# multitalk_audio_features.append(audio_feat.cpu().detach())
# fallback
if len(multitalk_audio_features) == 0:
raise RuntimeError("No valid audio embeddings extracted, please check inputs")
+203 -12
View File
@@ -1838,6 +1838,8 @@ class WanVideoSampler:
log.info(f"sigmas: {sample_scheduler.sigmas}")
else:
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
log.info(f"timesteps: {timesteps}")
total_steps = steps
steps = len(timesteps)
@@ -2237,6 +2239,35 @@ class WanVideoSampler:
log.info(f"mtv_motion_rotary_emb: {motion_rotary_emb[0].shape}")
mtv_freqs = mtv_freqs.to(device, dtype)
#region S2V
s2v_audio_input = s2v_ref_latent = s2v_pose = s2v_ref_motion = None
framepack = False
s2v_audio_embeds = image_embeds.get("audio_embeds", None)
if s2v_audio_embeds is not None:
log.info(f"Using S2V audio embeddings")
framepack = s2v_audio_embeds.get("enable_framepack", False)
if framepack and context_options is not None:
raise ValueError("S2V framepack and context windows cannot be used at the same time")
s2v_audio_input = s2v_audio_embeds.get("audio_embed_bucket", None)
if s2v_audio_input is not None:
#s2v_audio_input = s2v_audio_input[..., 0:image_embeds["num_frames"]]
s2v_audio_input = s2v_audio_input.to(device, dtype)
s2v_audio_scale = s2v_audio_embeds["audio_scale"]
s2v_ref_latent = s2v_audio_embeds.get("ref_latent", None)
if s2v_ref_latent is not None:
s2v_ref_latent = s2v_ref_latent.to(device, dtype)
s2v_ref_motion = s2v_audio_embeds.get("ref_motion", None)
if s2v_ref_motion is not None:
s2v_ref_motion = s2v_ref_motion.to(device, dtype)
s2v_pose = s2v_audio_embeds.get("pose_latent", None)
if s2v_pose is not None:
s2v_pose = s2v_pose.to(device, dtype)
s2v_pose_start_percent = s2v_audio_embeds.get("pose_start_percent", 0.0)
s2v_pose_end_percent = s2v_audio_embeds.get("pose_end_percent", 1.0)
s2v_num_repeat = s2v_audio_embeds.get("num_repeat", 1)
vae = s2v_audio_embeds.get("vae", None)
# vid2vid
noise_mask=original_image=None
if samples is not None and not multitalk_sampling:
@@ -2516,7 +2547,7 @@ class WanVideoSampler:
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False,
mtv_motion_tokens=None):
mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None):
nonlocal transformer
z = z.to(dtype)
autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear)
@@ -2661,7 +2692,11 @@ class WanVideoSampler:
else:
pcd_data_input = pcd_data
if s2v_pose is not None:
if not ((s2v_pose_start_percent <= current_step_percentage <= s2v_pose_end_percent) or \
(s2v_pose_end_percent > 0 and idx == 0 and current_step_percentage >= s2v_pose_start_percent)):
s2v_pose = None
base_params = {
'seq_len': seq_len, # sequence length
'device': device, # main device
@@ -2694,6 +2729,12 @@ 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 audio input
"s2v_ref_latent": s2v_ref_latent, # speech-to-video reference latent
"s2v_ref_motion": s2v_ref_motion, # speech-to-video reference motion latent
"s2v_audio_scale": s2v_audio_scale if s2v_audio_input is not None else 1.0, # speech-to-video audio scale
"s2v_pose": s2v_pose if s2v_pose is not None else None, # speech-to-video pose control
"s2v_motion_frames": s2v_motion_frames, # speech-to-video motion frames
}
batch_size = 1
@@ -2838,6 +2879,7 @@ class WanVideoSampler:
noise_pred = noise_pred_uncond_scaled + cfg_scale * filtered_cond * alpha
else:
noise_pred = noise_pred_uncond_scaled + cfg_scale * (noise_pred_cond - noise_pred_uncond_scaled)
del noise_pred_uncond_scaled, noise_pred_cond, noise_pred_uncond
return noise_pred, [cache_state_cond, cache_state_uncond]
@@ -2848,7 +2890,7 @@ class WanVideoSampler:
from .latent_preview import prepare_callback #custom for tiny VAE previews
callback = prepare_callback(patcher, len(timesteps))
if not multitalk_sampling:
if not multitalk_sampling and not framepack:
log.info(f"Input sequence length: {seq_len}")
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
@@ -3206,6 +3248,19 @@ class WanVideoSampler:
log.info(f"context window: {c}")
log.info(f"motion_token_indices: {start_token_index}-{end_token_index}")
partial_s2v_audio_input = None
if s2v_audio_input is not None:
indices = (torch.arange(4 + 1) - 2) * 1
audio_start = c[0] * 4
audio_end = c[-1] * 4 + 1
center_indices = torch.arange(audio_start, audio_end, 1)
center_indices = torch.clamp(center_indices, min=0, max=s2v_audio_input.shape[-1] - 1)
partial_s2v_audio_input = s2v_audio_input[..., center_indices]
partial_s2v_pose = None
if s2v_pose is not None:
partial_s2v_pose = s2v_pose[:, :, c].to(device, dtype)
partial_add_cond = None
if add_cond is not None:
partial_add_cond = add_cond[:, :, c].to(device, dtype)
@@ -3223,7 +3278,7 @@ class WanVideoSampler:
text_embeds["negative_prompt_embeds"],
partial_timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c, fantasy_portrait_input=partial_fantasy_portrait_input,
mtv_motion_tokens=partial_mtv_motion_tokens)
mtv_motion_tokens=partial_mtv_motion_tokens, s2v_audio_input=partial_s2v_audio_input, s2v_motion_frames=[1, 0], s2v_pose=partial_s2v_pose)
if cache_args is not None:
self.window_tracker.cache_states[window_id] = new_teacache
@@ -3231,7 +3286,7 @@ class WanVideoSampler:
window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped, window_type=context_options["fuse_method"])
noise_pred[:, c] += noise_pred_context * window_mask
counter[:, c] += window_mask
context_pbar.update_absolute(step_start_progress + (i + 1) * fraction_per_context, steps)
context_pbar.update_absolute(step_start_progress + (i + 1) * fraction_per_context, len(timesteps))
noise_pred /= counter
#region multitalk
elif multitalk_sampling:
@@ -3530,7 +3585,7 @@ class WanVideoSampler:
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
noise_pred, self.cache_state = predict_with_cfg(
latent_model_input, cfg[i], positive, text_embeds["negative_prompt_embeds"],
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input)
@@ -3648,21 +3703,157 @@ class WanVideoSampler:
except:
pass
return {"video": gen_video_samples.permute(1, 2, 3, 0)},
# region framepack loop
elif framepack:
framepack_out = []
ref_motion_image = None
#infer_frames = image_embeds["num_frames"]
infer_frames = s2v_audio_embeds.get("frame_window_size", 80)
motion_frames = infer_frames - 7 #73 default
lat_motion_frames = (motion_frames + 3) // 4
lat_target_frames = (infer_frames + 3 + motion_frames) // 4 - lat_motion_frames
step_iteration_count = 0
total_frames = s2v_audio_input.shape[-1]
s2v_motion_frames = [motion_frames, lat_motion_frames]
noise = torch.randn( #C, T, H, W
48 if is_5b else 16,
lat_target_frames,
target_shape[2],
target_shape[3],
dtype=torch.float32,
generator=seed_g,
device=torch.device("cpu"))
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
if ref_motion_image is None:
ref_motion_image = torch.zeros(
[1, 3, motion_frames, latent.shape[2]*vae_upscale_factor, latent.shape[3]*vae_upscale_factor],
dtype=vae.dtype,
device=device)
videos_last_frames = ref_motion_image
if s2v_pose is not None:
pose_cond_list = []
for r in range(s2v_num_repeat):
pose_start = r * (infer_frames // 4)
pose_end = pose_start + (infer_frames // 4)
cond_lat = s2v_pose[:, :, pose_start:pose_end]
pad_len = (infer_frames // 4) - cond_lat.shape[2]
if pad_len > 0:
pad = -torch.ones(cond_lat.shape[0], cond_lat.shape[1], pad_len, cond_lat.shape[3], cond_lat.shape[4], device=cond_lat.device, dtype=cond_lat.dtype)
cond_lat = torch.cat([cond_lat, pad], dim=2)
pose_cond_list.append(cond_lat.cpu())
log.info(f"Sampling {total_frames} frames in {s2v_num_repeat} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
# sample
for r in range(s2v_num_repeat):
vae.model.clear_cache()
mm.soft_empty_cache()
gc.collect()
if ref_motion_image is not None:
vae.to(device)
ref_motion = vae.encode(ref_motion_image.to(vae.dtype), device=device, pbar=False).to(dtype)[0]
vae.model.clear_cache()
vae.to(offload_device)
left_idx = r * infer_frames
right_idx = r * infer_frames + infer_frames
s2v_audio_input_slice = s2v_audio_input[..., left_idx:right_idx]
if s2v_audio_input_slice.shape[-1] < (right_idx - left_idx):
pad_len = (right_idx - left_idx) - s2v_audio_input_slice.shape[-1]
pad_shape = list(s2v_audio_input_slice.shape)
pad_shape[-1] = pad_len
pad = torch.zeros(pad_shape, device=s2v_audio_input_slice.device, dtype=s2v_audio_input_slice.dtype)
log.info(f"Padding s2v_audio_input_slice from {s2v_audio_input_slice.shape[-1]} to {right_idx - left_idx}")
s2v_audio_input_slice = torch.cat([s2v_audio_input_slice, pad], dim=-1)
if ref_motion_image is not None:
input_motion_latents = ref_motion.clone().unsqueeze(0)
else:
input_motion_latents = None
s2v_pose_slice = None
if s2v_pose is not None:
s2v_pose_slice = pose_cond_list[r].to(device)
sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
latent = noise.to(device)
for i, t in enumerate(tqdm(timesteps, desc=f"Sampling audio indices {left_idx}-{right_idx}", position=0)):
latent_model_input = latent.to(device)
timestep = torch.tensor([t]).to(device)
noise_pred, self.cache_state = predict_with_cfg(
latent_model_input,
cfg[idx],
text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens,
s2v_audio_input=s2v_audio_input_slice, s2v_ref_motion=input_motion_latents, s2v_motion_frames=s2v_motion_frames, s2v_pose=s2v_pose_slice)
latent = sample_scheduler.step(
noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0),
**scheduler_step_args)[0].squeeze(0)
if callback is not None:
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
callback(step_iteration_count, callback_latent, None, s2v_num_repeat*(len(timesteps)))
del callback_latent
step_iteration_count += 1
del latent_model_input, noise_pred
vae.to(device)
decode_latents = torch.cat([ref_motion.unsqueeze(0), latent.unsqueeze(0)], dim=2)
image = vae.decode(decode_latents.to(device, vae.dtype), device=device, pbar=False)[0]
del decode_latents
image = image.unsqueeze(0)[:, :, -infer_frames:]
if r == 0:
image = image[:, :, 3:]
framepack_out.append(image.cpu())
overlap_frames_num = min(motion_frames, image.shape[2])
videos_last_frames = torch.cat([
videos_last_frames[:, :, overlap_frames_num:],
image[:, :, -overlap_frames_num:]], dim=2).to(device, vae.dtype)
ref_motion_image = videos_last_frames
vae.to(offload_device)
vae.model.clear_cache()
mm.soft_empty_cache()
gen_video_samples = torch.cat(framepack_out, dim=2).squeeze(0).permute(1, 2, 3, 0)
if force_offload:
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
try:
print_memory(device)
torch.cuda.reset_peak_memory_stats(device)
except:
pass
return {"video": gen_video_samples},
#region normal inference
else:
noise_pred, self.cache_state = predict_with_cfg(
latent_model_input,
cfg[idx],
text_embeds["prompt_embeds"],
cfg[idx], text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens)
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input)
if bidirectional_sampling:
noise_pred_flipped, self.cache_state = predict_with_cfg(
latent_model_input_flipped,
cfg[idx],
text_embeds["prompt_embeds"],
cfg[idx], text_embeds["prompt_embeds"],
text_embeds["negative_prompt_embeds"],
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens,reverse_time=True)
+10 -3
View File
@@ -731,7 +731,8 @@ class WanVideoSetLoRAs:
def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
transformer_load_device=None, block_swap_args=None, gguf=False, reader=None, patcher=None):
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio_proj"}
params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding",
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer"}
param_count = sum(1 for _ in transformer.named_parameters())
pbar = ProgressBar(param_count)
cnt = 0
@@ -1081,7 +1082,9 @@ class WanVideoModelLoader:
ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1]
model_type = "t2v"
if not "text_embedding.0.weight" in sd:
if "audio_injector.injector.0.k.weight" in sd:
model_type = "s2v"
elif not "text_embedding.0.weight" in sd:
model_type = "no_cross_attn" #minimaxremover
elif "model_type.Wan2_1-FLF2V-14B-720P" in sd or "img_emb.emb_pos" in sd or "flf2v" in model.lower():
model_type = "fl2v"
@@ -1188,7 +1191,11 @@ class WanVideoModelLoader:
"add_ref_conv": True if "ref_conv.weight" in sd else False,
"in_dim_ref_conv": sd["ref_conv.weight"].shape[1] if "ref_conv.weight" in sd else None,
"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
"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,
"zero_timestep": model_type == "s2v",
}
with init_empty_weights():
+44 -2
View File
@@ -1,6 +1,7 @@
import torch
import numpy as np
from comfy.utils import common_upscale
from .utils import log
try:
from server import PromptServer
@@ -408,6 +409,45 @@ class WanVideoSigmaToStep:
def convert(self, sigma):
return (sigma,)
class NormalizeAudioLoudness:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"audio": ("AUDIO",),
"lufs": ("FLOAT", {"default": -23.0, "min": -100.0, "max": 0.0, "step": 0.1, "tool": "Loudness Units relative to Full Scale, higher LUFS values (closer to 0) mean louder audio. Lower LUFS values (more negative) mean quieter audio."}),
},
}
RETURN_TYPES = ("AUDIO", )
RETURN_NAMES = ("audio", )
FUNCTION = "normalize"
CATEGORY = "WanVideoWrapper"
def normalize(self, audio, lufs):
audio_input = audio["waveform"]
sample_rate = audio["sample_rate"]
if audio_input.dim() == 3:
audio_input = audio_input.squeeze(0)
audio_input_np = audio_input.detach().transpose(0, 1).numpy().astype(np.float32)
audio_input_np = np.ascontiguousarray(audio_input_np)
normalized_audio = self.loudness_norm(audio_input_np, sr=sample_rate, lufs=lufs)
out_audio = {"waveform": torch.from_numpy(normalized_audio).transpose(0, 1).unsqueeze(0).float(), "sample_rate": sample_rate}
return (out_audio, )
def loudness_norm(self, audio_array, sr=16000, lufs=-23):
try:
import pyloudnorm
except:
raise ImportError("pyloudnorm package is not installed")
meter = pyloudnorm.Meter(sr)
loudness = meter.integrated_loudness(audio_array)
if abs(loudness) > 100:
return audio_array
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
return normalized_audio
NODE_CLASS_MAPPINGS = {
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
@@ -416,7 +456,8 @@ NODE_CLASS_MAPPINGS = {
"DummyComfyWanModelObject": DummyComfyWanModelObject,
"WanVideoLatentReScale": WanVideoLatentReScale,
"CreateScheduleFloatList": CreateScheduleFloatList,
"WanVideoSigmaToStep": WanVideoSigmaToStep
"WanVideoSigmaToStep": WanVideoSigmaToStep,
"NormalizeAudioLoudness": NormalizeAudioLoudness
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
@@ -426,5 +467,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"DummyComfyWanModelObject": "Dummy Comfy Wan Model Object",
"WanVideoLatentReScale": "WanVideo Latent ReScale",
"CreateScheduleFloatList": "Create Schedule Float List",
"WanVideoSigmaToStep": "WanVideo Sigma To Step"
"WanVideoSigmaToStep": "WanVideo Sigma To Step",
"NormalizeAudioLoudness": "Normalize Audio Loudness"
}
+184
View File
@@ -0,0 +1,184 @@
import folder_paths
import math
import torch
import torch.nn.functional as F
import numpy as np
def get_sample_indices(original_fps,
total_frames,
target_fps,
num_sample,
fixed_start=None):
required_duration = num_sample / target_fps
required_origin_frames = int(np.ceil(required_duration * original_fps))
if required_duration > total_frames / original_fps:
raise ValueError("required_duration must be less than video length")
if not fixed_start is None and fixed_start >= 0:
start_frame = fixed_start
else:
max_start = total_frames - required_origin_frames
if max_start < 0:
raise ValueError("video length is too short")
start_frame = np.random.randint(0, max_start + 1)
start_time = start_frame / original_fps
end_time = start_time + required_duration
time_points = np.linspace(start_time, end_time, num_sample, endpoint=False)
frame_indices = np.round(np.array(time_points) * original_fps).astype(int)
frame_indices = np.clip(frame_indices, 0, total_frames - 1)
return frame_indices
def linear_interpolation(features, input_fps, output_fps, output_len=None):
"""
features: shape=[1, T, 512]
input_fps: fps for audio, f_a
output_fps: fps for video, f_m
output_len: video length
"""
features = features.transpose(1, 2) # [1, 512, T]
seq_len = features.shape[2] / float(input_fps) # T/f_a
if output_len is None:
output_len = int(seq_len * output_fps) # f_m*T/f_a
output_features = F.interpolate(
features, size=output_len, align_corners=True,
mode='linear') # [1, 512, output_len]
return output_features.transpose(1, 2) # [1, output_len, 512]
class WanVideoAddS2VEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"frame_window_size": ("INT", {"default": 80, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames in a single window"}),
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for audio embeddings"}),
"pose_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage for pose embeddings"}),
"pose_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage for pose embeddings"})
},
"optional": {
"audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",),
"ref_latent": ("LATENT",),
"pose_latent": ("LATENT",),
"vae": ("WANVAE",),
"enable_framepack": ("BOOLEAN", {"default": False, "tooltip": "Enable Framepack sampling loop, not compatible with context windows"})
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "INT",)
RETURN_NAMES = ("image_embeds", "audio_frame_count")
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, frame_window_size, audio_encoder_output=None, audio_scale=1.0, ref_latent=None, pose_latent=None, vae=None, pose_start_percent=0.0, pose_end_percent=1.0, enable_framepack=False):
audio_frame_count=0
if audio_encoder_output is not 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 in", audio_feat.shape)
input_fps = 50 # determined by the model itself
output_fps = 30 # determined by the model itself
bucket_fps = 16 # target fps for the generation
if input_fps != output_fps:
audio_feat = linear_interpolation(audio_feat, input_fps=input_fps, output_fps=output_fps)
print("audio_feat after interpolation", audio_feat.shape)
audio_feat = audio_feat[:, :embeds["num_frames"] * output_fps // bucket_fps, :]
print("audio_feat after trim", audio_feat.shape)
self.video_rate = output_fps
audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps(
audio_feat,
fps=bucket_fps,
batch_frames=frame_window_size
)
print("audio_embed_bucket", audio_embed_bucket.shape)
audio_embed_bucket = audio_embed_bucket.unsqueeze(0)
if len(audio_embed_bucket.shape) == 3:
audio_embed_bucket = audio_embed_bucket.permute(0, 2, 1)
elif len(audio_embed_bucket.shape) == 4:
audio_embed_bucket = audio_embed_bucket.permute(0, 2, 3, 1)
audio_frame_count = audio_embed_bucket.shape[-1]
print("audio_embed_bucket", audio_embed_bucket.shape)
new_entry = {
"audio_embed_bucket": audio_embed_bucket if audio_encoder_output is not None else None,
"num_repeat": num_repeat if audio_encoder_output is not None else None,
"ref_latent": ref_latent["samples"] if ref_latent is not None else None,
"pose_latent": pose_latent["samples"] if pose_latent is not None else None,
"audio_scale": audio_scale,
"vae": vae,
"pose_start_percent": pose_start_percent,
"pose_end_percent": pose_end_percent,
"enable_framepack": enable_framepack,
"frame_window_size": frame_window_size
}
updated = dict(embeds)
updated["audio_embeds"] = new_entry
return (updated, audio_frame_count)
def get_audio_embed_bucket_fps(self, audio_embed, fps=16, batch_frames=81, m=0):
num_layers, audio_frame_num, audio_dim = audio_embed.shape
if num_layers > 1:
return_all_layers = True
else:
return_all_layers = False
scale = self.video_rate / fps
min_batch_num = int(audio_frame_num / (batch_frames * scale)) + 1
bucket_num = min_batch_num * batch_frames
padd_audio_num = math.ceil(min_batch_num * batch_frames / fps * self.video_rate) - audio_frame_num
batch_idx = get_sample_indices(
original_fps=self.video_rate,
total_frames=audio_frame_num + padd_audio_num,
target_fps=fps,
num_sample=bucket_num,
fixed_start=0)
batch_audio_eb = []
audio_sample_stride = int(self.video_rate / fps)
for bi in batch_idx:
if bi < audio_frame_num:
chosen_idx = list(
range(bi - m * audio_sample_stride,
bi + (m + 1) * audio_sample_stride,
audio_sample_stride))
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
chosen_idx = [
audio_frame_num - 1 if c >= audio_frame_num else c
for c in chosen_idx
]
if return_all_layers:
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
start_dim=-2, end_dim=-1)
else:
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
else:
frame_audio_embed = \
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
batch_audio_eb.append(frame_audio_embed)
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
dim=0)
return batch_audio_eb, min_batch_num
NODE_CLASS_MAPPINGS = {
"WanVideoAddS2VEmbeds": WanVideoAddS2VEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddS2VEmbeds": "WanVideo Add S2V Embeds",
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+66 -40
View File
@@ -19,6 +19,11 @@ import comfy.model_management as mm
from comfy.utils import load_torch_file, ProgressBar, common_upscale
from comfy.clip_vision import clip_preprocess, ClipVisionModel
from comfy.cli_args import args, LatentPreviewMethod
from ..nodes_model_loading import load_weights
from ..nodes import offload_transformer
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
script_directory = os.path.dirname(os.path.abspath(__file__))
@@ -142,33 +147,51 @@ class WanVideoDiffusionForcingSampler:
patcher = model
model = model.model
transformer = model.diffusion_model
dtype = model["dtype"]
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
dtype = model["base_dtype"]
weight_dtype = model["weight_dtype"]
fp8_matmul = model["fp8_matmul"]
gguf = model["gguf"]
gguf_reader = model["gguf_reader"]
control_lora = model["control_lora"]
transformer_options = patcher.model_options.get("transformer_options", None)
merge_loras = transformer_options["merge_loras"]
patch_linear = transformer_options.get("patch_linear", False)
block_swap_args = transformer_options.get("block_swap_args", None)
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
transformer.blocks_to_swap = block_swap_args.get("blocks_to_swap", 0)
transformer.vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", 0)
transformer.prefetch_blocks = block_swap_args.get("prefetch_blocks", 0)
transformer.block_swap_debug = block_swap_args.get("block_swap_debug", False)
transformer.offload_img_emb = block_swap_args.get("offload_img_emb", False)
transformer.offload_txt_emb = block_swap_args.get("offload_txt_emb", False)
if gguf:
is_5b = transformer.out_dim == 48
vae_upscale_factor = 16 if is_5b else 8
# Load weights
if transformer.patched_linear and gguf_reader is None:
load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args)
if gguf_reader is not None: #handle GGUF
load_weights(transformer, patcher.model["sd"], base_dtype=dtype, transformer_load_device=device, patcher=patcher, gguf=True, reader=gguf_reader, block_swap_args=block_swap_args)
set_lora_params_gguf(transformer, patcher.patches)
elif len(patcher.patches) != 0 and patch_linear:
transformer.patched_linear = True
elif len(patcher.patches) != 0 and transformer.patched_linear: #handle patched linear layers (unmerged loras, fp8 scaled)
log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model")
if not merge_loras and fp8_matmul:
raise NotImplementedError("FP8 matmul with unmerged LoRAs is not supported")
set_lora_params(transformer, patcher.patches)
else:
remove_lora_from_module(transformer)
remove_lora_from_module(transformer) #clear possible unmerged lora weights
transformer.lora_scheduling_enabled = transformer_options.get("lora_scheduling_enabled", False)
#torch.compile
if model["auto_cpu_offload"] is False:
transformer = compile_model(transformer, model["compile_args"])
steps = int(steps/denoise_strength)
timesteps = None
@@ -367,34 +390,39 @@ class WanVideoDiffusionForcingSampler:
callback = prepare_callback(patcher, steps)
#blockswap init
if transformer_options is not None:
block_swap_args = transformer_options.get("block_swap_args", None)
#blockswap init
if not transformer.patched_linear:
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
if "control_adapter" in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
if block_swap_args is not None:
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", False)
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
elif model["manual_offloading"]:
transformer.to(device)
transformer.block_swap(
block_swap_args["blocks_to_swap"] - 1 ,
block_swap_args["offload_txt_emb"],
block_swap_args["offload_img_emb"],
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
prefetch_blocks = block_swap_args.get("prefetch_blocks", 0),
block_swap_debug = block_swap_args.get("block_swap_debug", False),
)
elif model["auto_cpu_offload"]:
for module in transformer.modules():
if hasattr(module, "offload"):
module.offload()
if hasattr(module, "onload"):
module.onload()
for block in transformer.blocks:
block.modulation = torch.nn.Parameter(block.modulation.to(device))
transformer.head.modulation = torch.nn.Parameter(transformer.head.modulation.to(device))
else:
transformer.to(device)
# Initialize Cache if enabled
transformer.enable_teacache = transformer.enable_magcache = False
@@ -610,10 +638,8 @@ class WanVideoDiffusionForcingSampler:
transformer.teacache_state.clear_all()
if force_offload:
if model["manual_offloading"]:
transformer.to(offload_device)
mm.soft_empty_cache()
gc.collect()
if not model["auto_cpu_offload"]:
offload_transformer(transformer)
try:
print_memory(device)
+503 -110
View File
@@ -20,7 +20,7 @@ except:
from .attention import attention
import numpy as np
from copy import deepcopy
from tqdm import tqdm
import gc
@@ -31,11 +31,94 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot
from ...MTV.mtv import apply_rotary_emb
#from .s2v.motioner import MotionerTransformers, FramePackMotioner, rope_precompute
#from comfy.ldm.wan.model import FramePackMotioner
class FramePackMotioner(nn.Module):
def __init__(
self,
inner_dim=1024,
num_heads=16, # Used to indicate the number of heads in the backbone network; unrelated to this module's design
zip_frame_buckets=[1, 2, 16], # Three numbers representing the number of frames sampled for patch operations from the nearest to the farthest frames
drop_mode="drop", # If not "drop", it will use "padd", meaning padding instead of deletion
):
super().__init__()
self.proj = nn.Conv3d(16, inner_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2))
self.proj_2x = nn.Conv3d(16, inner_dim, kernel_size=(2, 4, 4), stride=(2, 4, 4))
self.proj_4x = nn.Conv3d(16, inner_dim, kernel_size=(4, 8, 8), stride=(4, 8, 8))
self.zip_frame_buckets = zip_frame_buckets
self.inner_dim = inner_dim
self.num_heads = num_heads
self.drop_mode = drop_mode
def forward(self, motion_latents, rope_embedder, add_last_motion=2):
lat_height, lat_width = motion_latents.shape[3], motion_latents.shape[4]
padd_lat = torch.zeros(motion_latents.shape[0], 16, sum(self.zip_frame_buckets), lat_height, lat_width).to(device=motion_latents.device, dtype=motion_latents.dtype)
overlap_frame = min(padd_lat.shape[2], motion_latents.shape[2])
if overlap_frame > 0:
padd_lat[:, :, -overlap_frame:] = motion_latents[:, :, -overlap_frame:]
if add_last_motion < 2 and self.drop_mode != "drop":
zero_end_frame = sum(self.zip_frame_buckets[:len(self.zip_frame_buckets) - add_last_motion - 1])
padd_lat[:, :, -zero_end_frame:] = 0
clean_latents_4x, clean_latents_2x, clean_latents_post = padd_lat[:, :, -sum(self.zip_frame_buckets):, :, :].split(self.zip_frame_buckets[::-1], dim=2) # 16, 2 ,1
# patchfy
clean_latents_post = self.proj(clean_latents_post).flatten(2).transpose(1, 2)
clean_latents_2x = self.proj_2x(clean_latents_2x)
l_2x_shape = clean_latents_2x.shape
clean_latents_2x = clean_latents_2x.flatten(2).transpose(1, 2)
clean_latents_4x = self.proj_4x(clean_latents_4x)
l_4x_shape = clean_latents_4x.shape
clean_latents_4x = clean_latents_4x.flatten(2).transpose(1, 2)
if add_last_motion < 2 and self.drop_mode == "drop":
clean_latents_post = clean_latents_post[:, :0] if add_last_motion < 2 else clean_latents_post
clean_latents_2x = clean_latents_2x[:, :0] if add_last_motion < 1 else clean_latents_2x
motion_lat = torch.cat([clean_latents_post, clean_latents_2x, clean_latents_4x], dim=1)
rope_post = rope_embedder.rope_encode_comfy(1, lat_height, lat_width, t_start=-1, device=motion_latents.device, dtype=motion_latents.dtype)
rope_2x = rope_embedder.rope_encode_comfy(1, lat_height, lat_width, t_start=-3, steps_h=l_2x_shape[-2], steps_w=l_2x_shape[-1], device=motion_latents.device, dtype=motion_latents.dtype)
rope_4x = rope_embedder.rope_encode_comfy(4, lat_height, lat_width, t_start=-19, steps_h=l_4x_shape[-2], steps_w=l_4x_shape[-1], device=motion_latents.device, dtype=motion_latents.dtype)
rope = torch.cat([rope_post, rope_2x, rope_4x], dim=1)
return motion_lat, rope
from diffusers.models.attention import AdaLayerNorm
__all__ = ['WanModel']
from comfy import model_management as mm
def zero_module(module):
"""
Zero out the parameters of a module and return it.
"""
for p in module.parameters():
p.detach().zero_()
return module
def torch_dfs(model: nn.Module, parent_name='root'):
module_names, modules = [], []
current_name = parent_name if parent_name else 'root'
module_names.append(current_name)
modules.append(model)
for name, child in model.named_children():
if parent_name:
child_name = f'{parent_name}.{name}'
else:
child_name = name
child_modules, child_names = torch_dfs(child, child_name)
module_names += child_names
modules += child_modules
return modules, module_names
#from comfy.ldm.flux.math import apply_rope as apply_rope_comfy
def apply_rope_comfy(xq, xk, freqs_cis):
xq_ = xq.to(dtype=freqs_cis.dtype).reshape(*xq.shape[:-1], -1, 1, 2)
@@ -165,7 +248,6 @@ def sinusoidal_embedding_1d(dim, position):
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
return x
def rope_params(max_seq_len, dim, theta=10000, L_test=25, k=0):
assert dim % 2 == 0
exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim)
@@ -545,7 +627,7 @@ class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0,
num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
inner_t=None, inner_c=None, cross_freqs=None,
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, **kwargs):
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **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)
@@ -583,19 +665,20 @@ class WanT2VCrossAttention(WanSelfAttention):
# FantasyPortrait adapter attention
if adapter_proj is not None:
if len(adapter_proj.shape) == 4:
adapter_q = q.view(b * num_latent_frames, -1, n, d)
q_in = q[:, :orig_seq_len]
adapter_q = q_in.view(b * num_latent_frames, -1, n, d)
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
adapter_x = attention(adapter_q, ip_key, ip_value, attention_mode=self.attention_mode)
adapter_x = adapter_x.view(b, q.size(1), n, d)
adapter_x = adapter_x.view(b, q_in.size(1), n, d)
adapter_x = adapter_x.flatten(2)
elif len(adapter_proj.shape) == 3:
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b, -1, n, d)
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b, -1, n, d)
adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode)
adapter_x = attention(q_in, ip_key, ip_value, attention_mode=self.attention_mode)
adapter_x = adapter_x.flatten(2)
x = x + adapter_x * ip_scale
x[:, :orig_seq_len] = x[:, :orig_seq_len] + adapter_x * ip_scale
return self.o(x)
@@ -611,7 +694,7 @@ class WanI2VCrossAttention(WanSelfAttention):
def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None,
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy",
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, **kwargs):
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
r"""
Args:
x(Tensor): Shape [B, L1, C]
@@ -761,6 +844,7 @@ class WanAttentionBlock(nn.Module):
# modulation
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
self.seg_idx = None
@torch.compiler.disable()
def get_mod(self, e):
@@ -770,8 +854,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)
@@ -808,6 +908,7 @@ class WanAttentionBlock(nn.Module):
audio_proj=None,
audio_scale=1.0,
num_latent_frames=21,
original_seq_len=None,
enhance_enabled=False,
block_mask=None,
nag_params={},
@@ -838,9 +939,16 @@ 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.original_seq_len = original_seq_len
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))
@@ -958,7 +1066,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:
@@ -996,7 +1111,7 @@ class WanAttentionBlock(nn.Module):
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)
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=self.original_seq_len)
# 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,
@@ -1008,11 +1123,27 @@ class WanAttentionBlock(nn.Module):
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
if self.rope_func == "comfy_chunked":
if self.rope_func == "comfy_chunked" and not self.zero_timestep:
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()
@@ -1176,41 +1307,145 @@ class MLPProj(torch.nn.Module):
clip_extra_context_tokens = self.proj(image_embeds)
return clip_extra_context_tokens
from .s2v.auxi_blocks import MotionEncoder_tc
class CausalAudioEncoder(nn.Module):
def __init__(self,
dim=5120,
num_layers=25,
out_dim=2048,
video_rate=8,
num_token=4,
need_global=False):
super().__init__()
self.encoder = MotionEncoder_tc(
in_dim=dim,
hidden_dim=out_dim,
num_heads=num_token,
need_global=need_global)
weight = torch.ones((1, num_layers, 1, 1)) * 0.01
self.weights = torch.nn.Parameter(weight)
self.act = torch.nn.SiLU()
def forward(self, features):
# features B * num_layers * dim * video_length
weights = self.act(self.weights)
weights_sum = weights.sum(dim=1, keepdims=True)
weighted_feat = ((features * weights) / weights_sum).sum(
dim=1) # b dim f
weighted_feat = weighted_feat.permute(0, 2, 1) # b f dim
res = self.encoder(weighted_feat) # b f n dim
return res # b f n dim
class AudioCrossAttention(WanT2VCrossAttention):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
class AudioInjector_WAN(nn.Module):
def __init__(self,
all_modules,
all_modules_names,
dim=2048,
num_heads=32,
inject_layer=[0, 27],
root_net=None,
enable_adain=False,
adain_dim=2048,
need_adain_ont=False,
attention_mode='sdpa'):
super().__init__()
self.injected_block_id = {}
audio_injector_id = 0
for mod_name, mod in zip(all_modules_names, all_modules):
if isinstance(mod, WanAttentionBlock):
for inject_id in inject_layer:
if f'transformer_blocks.{inject_id}' in mod_name:
self.injected_block_id[inject_id] = audio_injector_id
audio_injector_id += 1
self.injector = nn.ModuleList([
AudioCrossAttention(
in_features=dim,
out_features=dim,
num_heads=num_heads,
qk_norm=True,
attention_mode=attention_mode
) for _ in range(audio_injector_id)
])
self.injector_pre_norm_feat = nn.ModuleList([
nn.LayerNorm(
dim,
elementwise_affine=False,
eps=1e-6,
) for _ in range(audio_injector_id)
])
self.injector_pre_norm_vec = nn.ModuleList([
nn.LayerNorm(
dim,
elementwise_affine=False,
eps=1e-6,
) for _ in range(audio_injector_id)
])
if enable_adain:
self.injector_adain_layers = nn.ModuleList([
AdaLayerNorm(
output_dim=dim * 2, embedding_dim=adain_dim, chunk_dim=1)
for _ in range(audio_injector_id)
])
if need_adain_ont:
self.injector_adain_output_layers = nn.ModuleList(
[nn.Linear(dim, dim) for _ in range(audio_injector_id)])
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',
main_device=torch.device('cuda'),
offload_device=torch.device('cpu'),
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
):
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',
main_device=torch.device('cuda'),
offload_device=torch.device('cpu'),
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",
audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
zero_timestep=False
):
r"""
Initialize the diffusion model backbone.
@@ -1361,7 +1596,7 @@ class WanModel(torch.nn.Module):
])
else:
# blocks
if model_type == 't2v':
if model_type == 't2v' or model_type == 's2v':
cross_attn_type = 't2v_cross_attn'
elif model_type == 'i2v' or model_type == 'fl2v':
cross_attn_type = 'i2v_cross_attn'
@@ -1416,6 +1651,47 @@ class WanModel(torch.nn.Module):
self.block_mask=None
#S2V
self.zero_timestep = self.audio_injector = self.trainable_cond_mask =None
if cond_dim > 0:
self.cond_encoder = nn.Conv3d(
cond_dim,
self.dim,
kernel_size=self.patch_size,
stride=self.patch_size)
if self.model_type == 's2v':
self.enable_adain = enable_adain
self.casual_audio_encoder = CausalAudioEncoder(
dim=audio_dim,
out_dim=self.dim,
num_token=num_audio_token,
need_global=enable_adain)
all_modules, all_modules_names = torch_dfs(
self.blocks, parent_name="root.transformer_blocks")
self.audio_injector = AudioInjector_WAN(
all_modules,
all_modules_names,
dim=self.dim,
num_heads=self.num_heads,
inject_layer=audio_inject_layers,
root_net=self,
enable_adain=enable_adain,
adain_dim=self.dim,
need_adain_ont=adain_mode != "attn_norm",
attention_mode=attention_mode
)
self.trainable_cond_mask = nn.Embedding(3, self.dim)
self.frame_packer = FramePackMotioner(
inner_dim=self.dim,
num_heads=self.num_heads,
zip_frame_buckets=[1, 2, 16],
drop_mode='padd')
self.adain_mode = adain_mode
self.zero_timestep = zero_timestep
@staticmethod
def _prepare_blockwise_causal_attn_mask(
device: torch.device | str, num_frames: int = 21,
@@ -1566,6 +1842,77 @@ class WanModel(torch.nn.Module):
block.to(self.offload_device, non_blocking=self.use_non_blocking)
return hints
def audio_injector_forward(self, block_idx, x, audio_emb, scale=1.0):
if block_idx in self.audio_injector.injected_block_id.keys():
audio_attn_id = self.audio_injector.injected_block_id[block_idx]
num_frames = audio_emb.shape[1]# b f n c
input_x = x[:, :self.original_seq_len].clone() # b (f h w) c
input_x = rearrange(input_x, "b (t n) c -> (b t) n c", t=num_frames)
if self.enable_adain and self.adain_mode == "attn_norm":
audio_emb_global = self.audio_emb_global
audio_emb_global = rearrange(audio_emb_global,"b t n c -> (b t) n c")
attn_x = self.audio_injector.injector_adain_layers[audio_attn_id](input_x, temb=audio_emb_global[:, 0])
else:
attn_x = self.audio_injector.injector_pre_norm_feat[audio_attn_id](input_x)
attn_audio_emb = rearrange(audio_emb, "b t n c -> (b t) n c", t=num_frames)
residual_out = self.audio_injector.injector[audio_attn_id](
x=attn_x ,
context=attn_audio_emb * scale,
)
residual_out = rearrange(residual_out, "(b t) n c -> b (t n) c", t=num_frames)
x[:, :self.original_seq_len].add_(residual_out)
return x
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
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 steps_t is None:
steps_t = t_len
if steps_h is None:
steps_h = h_len
if steps_w is None:
steps_w = w_len
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
if attn_cond is not None:
F_cond, H_cond, W_cond = attn_cond.shape[2], attn_cond.shape[3], attn_cond.shape[4]
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=device, dtype=dtype)
#shift
shift_f_size = 81 # Default value
shift_f = False
if shift_f:
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(shift_f_size, shift_f_size + cond_f_len - 1,steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
else:
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
# Combine original and conditional position ids
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
cond_img_ids = repeat(cond_img_ids, "t h w c -> b (t h w) c", b=1)
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
# Generate RoPE frequencies for the combined positions
freqs = self.rope_embedder(combined_img_ids, ntk_alphas).movedim(1, 2)
else:
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
return freqs
def forward(
self,
@@ -1611,7 +1958,13 @@ class WanModel(torch.nn.Module):
mtv_motion_rotary_emb=None,
mtv_freqs=None,
mtv_strength=1.0,
s2v_audio_input=None,
s2v_ref_latent=None,
s2v_audio_scale=1.0,
s2v_ref_motion=None,
s2v_pose=None,
s2v_motion_frames=[1, 0],
):
r"""
Forward pass through the diffusion model
@@ -1655,8 +2008,24 @@ class WanModel(torch.nn.Module):
if isinstance(submodule, nn.Linear):
if hasattr(submodule, 'step'):
submodule.step = current_step
#s2v
if self.model_type == 's2v' and s2v_audio_input is not None:
if is_uncond:
s2v_audio_input = s2v_audio_input * 0 # to match original code
s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, s2v_motion_frames[0]), s2v_audio_input], dim=-1)
audio_emb_res = self.casual_audio_encoder(s2v_audio_input)
if self.enable_adain:
audio_emb_global, audio_emb = audio_emb_res
self.audio_emb_global = audio_emb_global[:, s2v_motion_frames[1]:].clone()
else:
audio_emb = audio_emb_res
merged_audio_emb = audio_emb[:, s2v_motion_frames[1]:, :]
# params
device = self.patch_embedding.weight.device
if freqs is not None and freqs.device != device:
freqs = freqs.to(device)
@@ -1697,23 +2066,30 @@ class WanModel(torch.nn.Module):
for u in x
]
if s2v_pose is not None:
x[0] = x[0] + self.cond_encoder(s2v_pose.to(self.cond_encoder.weight.dtype)).to(x[0].dtype)
if self.control_adapter is not None and fun_camera is not None:
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])
original_grid_sizes = grid_sizes.clone()
x = [u.flatten(2).transpose(1, 2) for u in x]
x_len = x[0].shape[1]
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.float32)
assert seq_lens.max() <= seq_len
if self.trainable_cond_mask is not None:
cond_mask_weight = self.trainable_cond_mask.weight.to(x[0]).unsqueeze(1).unsqueeze(1)
self.original_seq_len = x[0].shape[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)
x[0] = x[0] + self.add_proj(add_cond)
if attn_cond is not None:
F_cond, H_cond, W_cond = attn_cond.shape[2], attn_cond.shape[3], attn_cond.shape[4]
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
attn_cond = self.attn_conv_in(attn_cond.to(self.attn_conv_in.weight.dtype)).to(x[0].dtype)
attn_cond = attn_cond.flatten(2).transpose(1, 2)
@@ -1727,24 +2103,34 @@ 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
F += end_ref_latent_frames
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)).to(x[0].dtype)
end_ref_latent = end_ref_latent.flatten(2).transpose(1, 2)
if cond_mask_weight is not None:
end_ref_latent = end_ref_latent + cond_mask_weight[1]
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
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.long)
assert seq_lens.max() <= seq_len
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
])
if self.trainable_cond_mask is not None:
x = x + cond_mask_weight[0]
# StandIn LoRA input
x_ip = None
freq_offset = 0
@@ -1761,10 +2147,9 @@ class WanModel(torch.nn.Module):
if freqs is None: #comfy rope
current_shape = (F, H, W)
has_cond = attn_cond is not None
f_len = ((F + (self.patch_size[0] // 2)) // self.patch_size[0])
h_len = ((H + (self.patch_size[1] // 2)) // self.patch_size[1])
w_len = ((W + (self.patch_size[2] // 2)) // self.patch_size[2])
if (self.cached_freqs is not None and
self.cached_shape == current_shape and
self.cached_cond == has_cond and
@@ -1773,37 +2158,14 @@ class WanModel(torch.nn.Module):
):
freqs = self.cached_freqs
else:
img_ids = torch.zeros((f_len, h_len, w_len, 3), device=x.device, dtype=x.dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(freq_offset, f_len + freq_offset - 1, steps=f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len + freq_offset - 1, steps=h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len + freq_offset - 1, steps=w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
if attn_cond is not None:
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=x.device, dtype=x.dtype)
#shift
shift_f_size = 81 # Default value
shift_f = False
if shift_f:
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(shift_f_size, shift_f_size + cond_f_len - 1,steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
else:
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
# Combine original and conditional position ids
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
cond_img_ids = repeat(cond_img_ids, "t h w c -> b (t h w) c", b=1)
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
# Generate RoPE frequencies for the combined positions
freqs = self.rope_embedder(combined_img_ids, ntk_alphas).movedim(1, 2)
else:
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
freqs = self.rope_encode_comfy(F, H, W, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond=attn_cond, device=x.device, dtype=x.dtype)
if s2v_ref_latent is not None:
freqs_ref = self.rope_encode_comfy(
s2v_ref_latent.shape[2],
s2v_ref_latent.shape[3],
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
@@ -1814,6 +2176,8 @@ class WanModel(torch.nn.Module):
# Stand-In RoPE frequencies
if x_ip is not None:
# Generate RoPE frequencies for x_ip
h_len = (H + 1) // 2
w_len = (W + 1) // 2
ip_img_ids = torch.zeros((f_ip, h_ip, w_ip, 3), device=x.device, dtype=x.dtype)
ip_img_ids[:, :, :, 0] = ip_img_ids[:, :, :, 0] + torch.linspace(0, f_ip - 1, steps=f_ip, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
ip_img_ids[:, :, :, 1] = ip_img_ids[:, :, :, 1] + torch.linspace(h_len + freq_offset, h_len + freq_offset + h_ip - 1, steps=h_ip, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
@@ -1828,6 +2192,15 @@ class WanModel(torch.nn.Module):
d = self.dim // self.num_heads
self.cross_freqs = rope_params(100, d).to(device=x.device)
if s2v_ref_motion is not None:
motion_encoded, freqs_motion = self.frame_packer(s2v_ref_motion, self)
motion_encoded = motion_encoded + cond_mask_weight[2]
x = torch.cat([x, motion_encoded], dim=1)
freqs = torch.cat([freqs, freqs_motion], dim=1)
#t = torch.repeat_interleave(t, 2, dim=1)
#t = torch.cat([t, torch.zeros((t.shape[0], 3), device=t.device, dtype=t.dtype)], dim=1)
# time embeddings
if t.dim() == 2:
b, f = t.shape
@@ -1835,9 +2208,23 @@ 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 )
@@ -2075,6 +2462,7 @@ class WanModel(torch.nn.Module):
camera_embed=camera_embed,
audio_proj=audio_proj,
num_latent_frames = F,
original_seq_len=self.original_seq_len,
enhance_enabled=enhance_enabled,
audio_scale=audio_scale,
block_mask=self.block_mask,
@@ -2171,6 +2559,8 @@ class WanModel(torch.nn.Module):
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
continue
x, x_ip = block(x, x_ip=x_ip, **kwargs) #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 self.block_swap_debug:
compute_end = time.perf_counter()
compute_time = compute_end - compute_start
@@ -2184,10 +2574,10 @@ class WanModel(torch.nn.Module):
#uni3c controlnet
if pdc_controlnet_states is not None and b < len(pdc_controlnet_states):
x[:, :x_len] += pdc_controlnet_states[b].to(x) * pcd_data["controlnet_weight"]
x[:, :self.original_seq_len] += pdc_controlnet_states[b].to(x) * pcd_data["controlnet_weight"]
#controlnet
if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])):
x[:, :x_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
x[:, :self.original_seq_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
if self.enable_teacache and (self.teacache_start_step <= current_step <= self.teacache_end_step) and pred_id is not None:
self.teacache_state.update(
@@ -2225,17 +2615,20 @@ 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]
x = x[:, :self.original_seq_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 = self.unpatchify(x, original_grid_sizes) # type: ignore[arg-type]
x = [u.float() for u in x]
return (x, pred_id) if pred_id is not None else (x, None)
+189
View File
@@ -0,0 +1,189 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import math
import librosa
import numpy as np
import torch
import torch.nn.functional as F
from transformers import Wav2Vec2ForCTC, Wav2Vec2Processor
def get_sample_indices(original_fps,
total_frames,
target_fps,
num_sample,
fixed_start=None):
required_duration = num_sample / target_fps
required_origin_frames = int(np.ceil(required_duration * original_fps))
if required_duration > total_frames / original_fps:
raise ValueError("required_duration must be less than video length")
if not fixed_start is None and fixed_start >= 0:
start_frame = fixed_start
else:
max_start = total_frames - required_origin_frames
if max_start < 0:
raise ValueError("video length is too short")
start_frame = np.random.randint(0, max_start + 1)
start_time = start_frame / original_fps
end_time = start_time + required_duration
time_points = np.linspace(start_time, end_time, num_sample, endpoint=False)
frame_indices = np.round(np.array(time_points) * original_fps).astype(int)
frame_indices = np.clip(frame_indices, 0, total_frames - 1)
return frame_indices
def linear_interpolation(features, input_fps, output_fps, output_len=None):
"""
features: shape=[1, T, 512]
input_fps: fps for audio, f_a
output_fps: fps for video, f_m
output_len: video length
"""
features = features.transpose(1, 2) # [1, 512, T]
seq_len = features.shape[2] / float(input_fps) # T/f_a
if output_len is None:
output_len = int(seq_len * output_fps) # f_m*T/f_a
output_features = F.interpolate(
features, size=output_len, align_corners=True,
mode='linear') # [1, 512, output_len]
return output_features.transpose(1, 2) # [1, output_len, 512]
class AudioEncoder():
def __init__(self, device='cpu', model_id="facebook/wav2vec2-base-960h"):
# load pretrained model
self.processor = Wav2Vec2Processor.from_pretrained(model_id)
self.model = Wav2Vec2ForCTC.from_pretrained(model_id)
self.model = self.model.to(device)
self.video_rate = 30
def extract_audio_feat(self,
audio_path,
return_all_layers=False,
dtype=torch.float32):
audio_input, sample_rate = librosa.load(audio_path, sr=16000)
input_values = self.processor(
audio_input, sampling_rate=sample_rate,
return_tensors="pt").input_values
# INFERENCE
# retrieve logits & take argmax
res = self.model(
input_values.to(self.model.device), output_hidden_states=True)
if return_all_layers:
feat = torch.cat(res.hidden_states)
else:
feat = res.hidden_states[-1]
feat = linear_interpolation(
feat, input_fps=50, output_fps=self.video_rate)
z = feat.to(dtype) # Encoding for the motion
return z
def get_audio_embed_bucket(self,
audio_embed,
stride=2,
batch_frames=12,
m=2):
num_layers, audio_frame_num, audio_dim = audio_embed.shape
if num_layers > 1:
return_all_layers = True
else:
return_all_layers = False
min_batch_num = int(audio_frame_num / (batch_frames * stride)) + 1
bucket_num = min_batch_num * batch_frames
batch_idx = [stride * i for i in range(bucket_num)]
batch_audio_eb = []
for bi in batch_idx:
if bi < audio_frame_num:
audio_sample_stride = 2
chosen_idx = list(
range(bi - m * audio_sample_stride,
bi + (m + 1) * audio_sample_stride,
audio_sample_stride))
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
chosen_idx = [
audio_frame_num - 1 if c >= audio_frame_num else c
for c in chosen_idx
]
if return_all_layers:
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
start_dim=-2, end_dim=-1)
else:
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
else:
frame_audio_embed = \
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
batch_audio_eb.append(frame_audio_embed)
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
dim=0)
return batch_audio_eb, min_batch_num
def get_audio_embed_bucket_fps(self,
audio_embed,
fps=16,
batch_frames=81,
m=0):
num_layers, audio_frame_num, audio_dim = audio_embed.shape
if num_layers > 1:
return_all_layers = True
else:
return_all_layers = False
scale = self.video_rate / fps
min_batch_num = int(audio_frame_num / (batch_frames * scale)) + 1
bucket_num = min_batch_num * batch_frames
padd_audio_num = math.ceil(min_batch_num * batch_frames / fps *
self.video_rate) - audio_frame_num
batch_idx = get_sample_indices(
original_fps=self.video_rate,
total_frames=audio_frame_num + padd_audio_num,
target_fps=fps,
num_sample=bucket_num,
fixed_start=0)
batch_audio_eb = []
audio_sample_stride = int(self.video_rate / fps)
for bi in batch_idx:
if bi < audio_frame_num:
chosen_idx = list(
range(bi - m * audio_sample_stride,
bi + (m + 1) * audio_sample_stride,
audio_sample_stride))
chosen_idx = [0 if c < 0 else c for c in chosen_idx]
chosen_idx = [
audio_frame_num - 1 if c >= audio_frame_num else c
for c in chosen_idx
]
if return_all_layers:
frame_audio_embed = audio_embed[:, chosen_idx].flatten(
start_dim=-2, end_dim=-1)
else:
frame_audio_embed = audio_embed[0][chosen_idx].flatten()
else:
frame_audio_embed = \
torch.zeros([audio_dim * (2 * m + 1)], device=audio_embed.device) if not return_all_layers \
else torch.zeros([num_layers, audio_dim * (2 * m + 1)], device=audio_embed.device)
batch_audio_eb.append(frame_audio_embed)
batch_audio_eb = torch.cat([c.unsqueeze(0) for c in batch_audio_eb],
dim=0)
return batch_audio_eb, min_batch_num
+129
View File
@@ -0,0 +1,129 @@
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
class CausalConv1d(nn.Module):
def __init__(self,
chan_in,
chan_out,
kernel_size=3,
stride=1,
dilation=1,
pad_mode='replicate',
**kwargs):
super().__init__()
self.pad_mode = pad_mode
padding = (kernel_size - 1, 0) # T
self.time_causal_padding = padding
self.conv = nn.Conv1d(
chan_in,
chan_out,
kernel_size,
stride=stride,
dilation=dilation,
**kwargs)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
return self.conv(x)
class MotionEncoder_tc(nn.Module):
def __init__(self,
in_dim: int,
hidden_dim: int,
num_heads=int,
need_global=True,
dtype=None,
device=None):
factory_kwargs = {"dtype": dtype, "device": device}
super().__init__()
self.num_heads = num_heads
self.need_global = need_global
self.conv1_local = CausalConv1d(
in_dim, hidden_dim // 4 * num_heads, 3, stride=1)
if need_global:
self.conv1_global = CausalConv1d(
in_dim, hidden_dim // 4, 3, stride=1)
self.norm1 = nn.LayerNorm(
hidden_dim // 4,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.act = nn.SiLU()
self.conv2 = CausalConv1d(hidden_dim // 4, hidden_dim // 2, 3, stride=2)
self.conv3 = CausalConv1d(hidden_dim // 2, hidden_dim, 3, stride=2)
if need_global:
self.final_linear = nn.Linear(hidden_dim, hidden_dim,
**factory_kwargs)
self.norm1 = nn.LayerNorm(
hidden_dim // 4,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.norm2 = nn.LayerNorm(
hidden_dim // 2,
elementwise_affine=False,
eps=1e-6,
**factory_kwargs)
self.norm3 = nn.LayerNorm(
hidden_dim, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.padding_tokens = nn.Parameter(torch.zeros(1, 1, 1, hidden_dim))
def forward(self, x):
x = rearrange(x, 'b t c -> b c t')
x_ori = x.clone()
b, c, t = x.shape
x = self.conv1_local(x)
x = rearrange(x, 'b (n c) t -> (b n) t c', n=self.num_heads)
x = self.norm1(x)
x = self.act(x)
x = rearrange(x, 'b t c -> b c t')
x = self.conv2(x)
x = rearrange(x, 'b c t -> b t c')
x = self.norm2(x)
x = self.act(x)
x = rearrange(x, 'b t c -> b c t')
x = self.conv3(x)
x = rearrange(x, 'b c t -> b t c')
x = self.norm3(x)
x = self.act(x)
x = rearrange(x, '(b n) t c -> b t n c', b=b)
padding = self.padding_tokens.repeat(b, x.shape[1], 1, 1)
x = torch.cat([x, padding], dim=-2)
x_local = x.clone()
if not self.need_global:
return x_local
x = self.conv1_global(x_ori)
x = rearrange(x, 'b c t -> b t c')
x = self.norm1(x)
x = self.act(x)
x = rearrange(x, 'b t c -> b c t')
x = self.conv2(x)
x = rearrange(x, 'b c t -> b t c')
x = self.norm2(x)
x = self.act(x)
x = rearrange(x, 'b t c -> b c t')
x = self.conv3(x)
x = rearrange(x, 'b c t -> b t c')
x = self.norm3(x)
x = self.act(x)
x = self.final_linear(x)
x = rearrange(x, '(b n) t c -> b t n c', b=b)
return x, x_local
+5 -4
View File
@@ -22,7 +22,7 @@ scheduler_list = [
"multitalk"
]
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None):
def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer_dim=5120, flowedit_args=None, denoise_strength=1.0, sigmas=None, log_timesteps=False):
timesteps = None
if 'unipc' in scheduler:
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
@@ -111,7 +111,8 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
start_idx = 0
end_idx = len(timesteps) - 1
log.info(f"Total timesteps: {timesteps}")
if log_timesteps:
log.info(f"Total timesteps: {timesteps}")
if isinstance(start_step, float):
idxs = (sample_scheduler.sigmas <= start_step).nonzero(as_tuple=True)[0]
@@ -134,8 +135,8 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo
sample_scheduler.full_sigmas = sample_scheduler.sigmas.clone()
sample_scheduler.sigmas = sample_scheduler.sigmas[start_idx:start_idx+len(timesteps)+1] # always one longer
log.info(f"Using timesteps: {timesteps}")
if log_timesteps:
log.info(f"Using timesteps: {timesteps}")
if hasattr(sample_scheduler, 'timesteps'):
sample_scheduler.timesteps = timesteps