From 748ec89aa8edf7b108006e936a457ff46ee84391 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 26 Aug 2025 19:02:24 +0300 Subject: [PATCH 01/31] init --- __init__.py | 3 + nodes.py | 11 + nodes_model_loading.py | 11 +- s2v/nodes.py | 164 ++++++ wanvideo/modules/model.py | 289 ++++++++-- wanvideo/modules/s2v/audio_encoder.py | 189 ++++++ wanvideo/modules/s2v/auxi_blocks.py | 129 +++++ wanvideo/modules/s2v/motioner.py | 794 ++++++++++++++++++++++++++ wanvideo/modules/s2v/s2v_utils.py | 70 +++ 9 files changed, 1623 insertions(+), 37 deletions(-) create mode 100644 s2v/nodes.py create mode 100644 wanvideo/modules/s2v/audio_encoder.py create mode 100644 wanvideo/modules/s2v/auxi_blocks.py create mode 100644 wanvideo/modules/s2v/motioner.py create mode 100644 wanvideo/modules/s2v/s2v_utils.py diff --git a/__init__.py b/__init__.py index 1f81508..8ea20c2 100644 --- a/__init__.py +++ b/__init__.py @@ -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"] \ No newline at end of file diff --git a/nodes.py b/nodes.py index 37e7ea4..1abe0a9 100644 --- a/nodes.py +++ b/nodes.py @@ -2212,6 +2212,16 @@ 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 = 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_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) + # vid2vid noise_mask=original_image=None if samples is not None and not multitalk_sampling: @@ -2669,6 +2679,7 @@ 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 } batch_size = 1 diff --git a/nodes_model_loading.py b/nodes_model_loading.py index a5f6567..cbabd33 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -731,7 +731,7 @@ 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", "audio"} param_count = sum(1 for _ in transformer.named_parameters()) pbar = ProgressBar(param_count) cnt = 0 @@ -1079,7 +1079,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" @@ -1186,7 +1188,10 @@ 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 + } with init_empty_weights(): diff --git a/s2v/nodes.py b/s2v/nodes.py new file mode 100644 index 0000000..bb4464e --- /dev/null +++ b/s2v/nodes.py @@ -0,0 +1,164 @@ +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 WanVideoAddAudioEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "embeds": ("WANVIDIMAGE_EMBEDS",), + "audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",), + "input_fps": ("FLOAT", {"default": 30.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"}), + "frames": ("INT", {"default": 81, "min": 1, "max": 120, "step": 1, "tooltip": "Number of frames to process"}) + } + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "add" + CATEGORY = "WanVideoWrapper" + + def add(self, embeds, input_fps, output_fps, frames, audio_encoder_output): + # Prepare the new audio entry + + #audio_feat = audio_encoder_output["encoded_audio"] + #print("audio_feat", audio_feat.shape) + 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) + + if input_fps != output_fps: + audio_feat = linear_interpolation(audio_feat, input_fps=input_fps, output_fps=output_fps) + + self.video_rate = output_fps + + audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps( + audio_feat, + fps=output_fps, + batch_frames=frames + ) + + 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_embed_bucket = audio_embed_bucket[..., 0:frames] + + print("audio_embed_bucket", audio_embed_bucket.shape) + + new_entry = { + "audio_embed_bucket": audio_embed_bucket, + "num_repeat": num_repeat + } + updated = dict(embeds) + updated["audio_embeds"] = new_entry + return (updated,) + + 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 = { + "WanVideoAddAudioEmbeds": WanVideoAddAudioEmbeds, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "WanVideoAddAudioEmbeds": "WanVideo Add Audio Embeds", +} \ No newline at end of file diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 8bb23ac..d232e89 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -30,11 +30,38 @@ from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot from ...MTV.mtv import apply_rotary_emb +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) @@ -1175,41 +1202,143 @@ 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): + super().__init__() + num_injector_layers = len(inject_layer) + 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, + ) 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], + ): r""" Initialize the diffusion model backbone. @@ -1360,7 +1489,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' @@ -1415,6 +1544,36 @@ class WanModel(torch.nn.Module): self.block_mask=None + #S2V + if cond_dim > 0: + self.cond_encoder = nn.Conv3d( + cond_dim, + self.dim, + kernel_size=self.patch_size, + stride=self.patch_size) + 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", + ) + self.adain_mode = adain_mode + + self.trainable_cond_mask = nn.Embedding(3, self.dim) + @staticmethod def _prepare_blockwise_causal_attn_mask( device: torch.device | str, num_frames: int = 21, @@ -1565,6 +1724,46 @@ 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, hidden_states, merged_audio_emb): + if block_idx in self.audio_injector.injected_block_id.keys(): + audio_attn_id = self.audio_injector.injected_block_id[block_idx] + audio_emb = merged_audio_emb # b f n c + num_frames = audio_emb.shape[1] + + input_hidden_states = hidden_states[:, :self.original_seq_len].clone() # b (f h w) c + input_hidden_states = rearrange( + input_hidden_states, "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") + adain_hidden_states = self.audio_injector.injector_adain_layers[ + audio_attn_id]( + input_hidden_states, temb=audio_emb_global[:, 0]) + attn_hidden_states = adain_hidden_states + else: + attn_hidden_states = self.audio_injector.injector_pre_norm_feat[ + audio_attn_id]( + input_hidden_states) + audio_emb = rearrange( + audio_emb, "b t n c -> (b t) n c", t=num_frames) + attn_audio_emb = audio_emb + residual_out = self.audio_injector.injector[audio_attn_id]( + x=attn_hidden_states, + context=attn_audio_emb, + context_lens=torch.ones( + attn_hidden_states.shape[0], + dtype=torch.long, + device=attn_hidden_states.device) * attn_audio_emb.shape[1]) + residual_out = rearrange( + residual_out, "(b t) n c -> b (t n) c", t=num_frames) + hidden_states[:, :self. + original_seq_len] = hidden_states[:, :self. + original_seq_len] + residual_out + + return hidden_states def forward( self, @@ -1610,6 +1809,7 @@ class WanModel(torch.nn.Module): mtv_motion_rotary_emb=None, mtv_freqs=None, mtv_strength=1.0, + s2v_audio_input=None ): r""" @@ -1654,8 +1854,25 @@ 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: + motion_frames=[17, 5] + + s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, 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[:, motion_frames[1]:].clone() + else: + audio_emb = audio_emb_res + merged_audio_emb = audio_emb[:, motion_frames[1]:, :] + + # params device = self.patch_embedding.weight.device + if freqs is not None and freqs.device != device: freqs = freqs.to(device) @@ -1737,7 +1954,9 @@ class WanModel(torch.nn.Module): F += phantom_ref_frames x = [torch.concat([u, phantom_ref.unsqueeze(0)], dim=1) for phantom_ref, u in zip(phantom_ref, x)] - seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long) + 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 x = torch.cat([ torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], @@ -2163,6 +2382,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) #s2v if self.block_swap_debug: compute_end = time.perf_counter() compute_time = compute_end - compute_start diff --git a/wanvideo/modules/s2v/audio_encoder.py b/wanvideo/modules/s2v/audio_encoder.py new file mode 100644 index 0000000..05fea4e --- /dev/null +++ b/wanvideo/modules/s2v/audio_encoder.py @@ -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 diff --git a/wanvideo/modules/s2v/auxi_blocks.py b/wanvideo/modules/s2v/auxi_blocks.py new file mode 100644 index 0000000..bef8f33 --- /dev/null +++ b/wanvideo/modules/s2v/auxi_blocks.py @@ -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 diff --git a/wanvideo/modules/s2v/motioner.py b/wanvideo/modules/s2v/motioner.py new file mode 100644 index 0000000..699c570 --- /dev/null +++ b/wanvideo/modules/s2v/motioner.py @@ -0,0 +1,794 @@ +# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. +import math +from typing import Any, Dict, List, Literal, Optional, Union + +import numpy as np +import torch +import torch.cuda.amp as amp +import torch.nn as nn +from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin +from diffusers.utils import BaseOutput, is_torch_version +from einops import rearrange, repeat + +from ..model import flash_attention +from .s2v_utils import rope_precompute + + +def sinusoidal_embedding_1d(dim, position): + # preprocess + assert dim % 2 == 0 + half = dim // 2 + position = position.type(torch.float64) + + # calculation + sinusoid = torch.outer( + position, torch.pow(10000, -torch.arange(half).to(position).div(half))) + x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1) + return x + + +@amp.autocast(enabled=False) +def rope_params(max_seq_len, dim, theta=10000): + assert dim % 2 == 0 + freqs = torch.outer( + torch.arange(max_seq_len), + 1.0 / torch.pow(theta, + torch.arange(0, dim, 2).to(torch.float64).div(dim))) + freqs = torch.polar(torch.ones_like(freqs), freqs) + return freqs + + +@amp.autocast(enabled=False) +def rope_apply(x, grid_sizes, freqs, start=None): + n, c = x.size(2), x.size(3) // 2 + + # split freqs + if type(freqs) is list: + trainable_freqs = freqs[1] + freqs = freqs[0] + freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1) + + # loop over samples + output = [] + output = x.clone() + seq_bucket = [0] + if not type(grid_sizes) is list: + grid_sizes = [grid_sizes] + for g in grid_sizes: + if not type(g) is list: + g = [torch.zeros_like(g), g] + batch_size = g[0].shape[0] + for i in range(batch_size): + if start is None: + f_o, h_o, w_o = g[0][i] + else: + f_o, h_o, w_o = start[i] + + f, h, w = g[1][i] + t_f, t_h, t_w = g[2][i] + seq_f, seq_h, seq_w = f - f_o, h - h_o, w - w_o + seq_len = int(seq_f * seq_h * seq_w) + if seq_len > 0: + if t_f > 0: + factor_f, factor_h, factor_w = (t_f / seq_f).item(), ( + t_h / seq_h).item(), (t_w / seq_w).item() + + if f_o >= 0: + f_sam = np.linspace(f_o.item(), (t_f + f_o).item() - 1, + seq_f).astype(int).tolist() + else: + f_sam = np.linspace(-f_o.item(), + (-t_f - f_o).item() + 1, + seq_f).astype(int).tolist() + h_sam = np.linspace(h_o.item(), (t_h + h_o).item() - 1, + seq_h).astype(int).tolist() + w_sam = np.linspace(w_o.item(), (t_w + w_o).item() - 1, + seq_w).astype(int).tolist() + + assert f_o * f >= 0 and h_o * h >= 0 and w_o * w >= 0 + freqs_0 = freqs[0][f_sam] if f_o >= 0 else freqs[0][ + f_sam].conj() + freqs_0 = freqs_0.view(seq_f, 1, 1, -1) + + freqs_i = torch.cat([ + freqs_0.expand(seq_f, seq_h, seq_w, -1), + freqs[1][h_sam].view(1, seq_h, 1, -1).expand( + seq_f, seq_h, seq_w, -1), + freqs[2][w_sam].view(1, 1, seq_w, -1).expand( + seq_f, seq_h, seq_w, -1), + ], + dim=-1).reshape(seq_len, 1, -1) + elif t_f < 0: + freqs_i = trainable_freqs.unsqueeze(1) + # apply rotary embedding + # precompute multipliers + x_i = torch.view_as_complex( + x[i, seq_bucket[-1]:seq_bucket[-1] + seq_len].to( + torch.float64).reshape(seq_len, n, -1, 2)) + x_i = torch.view_as_real(x_i * freqs_i).flatten(2) + output[i, seq_bucket[-1]:seq_bucket[-1] + seq_len] = x_i + seq_bucket.append(seq_bucket[-1] + seq_len) + return output.float() + + +class RMSNorm(nn.Module): + + def __init__(self, dim, eps=1e-5): + super().__init__() + self.dim = dim + self.eps = eps + self.weight = nn.Parameter(torch.ones(dim)) + + def forward(self, x): + return self._norm(x.float()).type_as(x) * self.weight + + def _norm(self, x): + return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) + + +class LayerNorm(nn.LayerNorm): + + def __init__(self, dim, eps=1e-6, elementwise_affine=False): + super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps) + + def forward(self, x): + return super().forward(x.float()).type_as(x) + + +class SelfAttention(nn.Module): + + def __init__(self, + dim, + num_heads, + window_size=(-1, -1), + qk_norm=True, + eps=1e-6): + assert dim % num_heads == 0 + super().__init__() + self.dim = dim + self.num_heads = num_heads + self.head_dim = dim // num_heads + self.window_size = window_size + self.qk_norm = qk_norm + self.eps = eps + + # layers + self.q = nn.Linear(dim, dim) + self.k = nn.Linear(dim, dim) + self.v = nn.Linear(dim, dim) + self.o = nn.Linear(dim, dim) + self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() + self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() + + def forward(self, x, seq_lens, grid_sizes, freqs): + b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim + + # query, key, value function + def qkv_fn(x): + q = self.norm_q(self.q(x)).view(b, s, n, d) + k = self.norm_k(self.k(x)).view(b, s, n, d) + v = self.v(x).view(b, s, n, d) + return q, k, v + + q, k, v = qkv_fn(x) + + x = flash_attention( + q=rope_apply(q, grid_sizes, freqs), + k=rope_apply(k, grid_sizes, freqs), + v=v, + k_lens=seq_lens, + window_size=self.window_size) + + # output + x = x.flatten(2) + x = self.o(x) + return x + + +class SwinSelfAttention(SelfAttention): + + def forward(self, x, seq_lens, grid_sizes, freqs): + b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim + assert b == 1, 'Only support batch_size 1' + + # query, key, value function + def qkv_fn(x): + q = self.norm_q(self.q(x)).view(b, s, n, d) + k = self.norm_k(self.k(x)).view(b, s, n, d) + v = self.v(x).view(b, s, n, d) + return q, k, v + + q, k, v = qkv_fn(x) + + q = rope_apply(q, grid_sizes, freqs) + k = rope_apply(k, grid_sizes, freqs) + T, H, W = grid_sizes[0].tolist() + + q = rearrange(q, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) + k = rearrange(k, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) + v = rearrange(v, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) + + ref_q = q[-1:] + q = q[:-1] + + ref_k = repeat( + k[-1:], "1 s n d -> t s n d", t=k.shape[0] - 1) # t hw n d + k = k[:-1] + k = torch.cat([k[:1], k, k[-1:]]) + k = torch.cat([k[1:-1], k[2:], k[:-2], ref_k], dim=1) # (bt) (3hw) n d + + ref_v = repeat(v[-1:], "1 s n d -> t s n d", t=v.shape[0] - 1) + v = v[:-1] + v = torch.cat([v[:1], v, v[-1:]]) + v = torch.cat([v[1:-1], v[2:], v[:-2], ref_v], dim=1) + + # q: b (t h w) n d + # k: b (t h w) n d + out = flash_attention( + q=q, + k=k, + v=v, + # k_lens=torch.tensor([k.shape[1]] * k.shape[0], device=x.device, dtype=torch.long), + window_size=self.window_size) + out = torch.cat([out, ref_v[:1]], axis=0) + out = rearrange(out, '(b t) (h w) n d -> b (t h w) n d', t=T, h=H, w=W) + x = out + + # output + x = x.flatten(2) + x = self.o(x) + return x + + +#Fix the reference frame RoPE to 1,H,W. +#Set the current frame RoPE to 1. +#Set the previous frame RoPE to 0. +class CasualSelfAttention(SelfAttention): + + def forward(self, x, seq_lens, grid_sizes, freqs): + shifting = 3 + b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim + assert b == 1, 'Only support batch_size 1' + + # query, key, value function + def qkv_fn(x): + q = self.norm_q(self.q(x)).view(b, s, n, d) + k = self.norm_k(self.k(x)).view(b, s, n, d) + v = self.v(x).view(b, s, n, d) + return q, k, v + + q, k, v = qkv_fn(x) + + T, H, W = grid_sizes[0].tolist() + + q = rearrange(q, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) + k = rearrange(k, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) + v = rearrange(v, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) + + ref_q = q[-1:] + q = q[:-1] + + grid_sizes = torch.tensor([[1, H, W]] * q.shape[0], dtype=torch.long) + start = [[shifting, 0, 0]] * q.shape[0] + q = rope_apply(q, grid_sizes, freqs, start=start) + + ref_k = k[-1:] + grid_sizes = torch.tensor([[1, H, W]], dtype=torch.long) + # start = [[shifting, H, W]] + + start = [[shifting + 10, 0, 0]] + ref_k = rope_apply(ref_k, grid_sizes, freqs, start) + ref_k = repeat( + ref_k, "1 s n d -> t s n d", t=k.shape[0] - 1) # t hw n d + + k = k[:-1] + k = torch.cat([*([k[:1]] * shifting), k]) + cat_k = [] + for i in range(shifting): + cat_k.append(k[i:i - shifting]) + cat_k.append(k[shifting:]) + k = torch.cat(cat_k, dim=1) # (bt) (3hw) n d + + grid_sizes = torch.tensor( + [[shifting + 1, H, W]] * q.shape[0], dtype=torch.long) + k = rope_apply(k, grid_sizes, freqs) + k = torch.cat([k, ref_k], dim=1) + + ref_v = repeat(v[-1:], "1 s n d -> t s n d", t=q.shape[0]) # t hw n d + v = v[:-1] + v = torch.cat([*([v[:1]] * shifting), v]) + cat_v = [] + for i in range(shifting): + cat_v.append(v[i:i - shifting]) + cat_v.append(v[shifting:]) + v = torch.cat(cat_v, dim=1) # (bt) (3hw) n d + v = torch.cat([v, ref_v], dim=1) + + # q: b (t h w) n d + # k: b (t h w) n d + outs = [] + for i in range(q.shape[0]): + out = flash_attention( + q=q[i:i + 1], + k=k[i:i + 1], + v=v[i:i + 1], + window_size=self.window_size) + outs.append(out) + out = torch.cat(outs, dim=0) + out = torch.cat([out, ref_v[:1]], axis=0) + out = rearrange(out, '(b t) (h w) n d -> b (t h w) n d', t=T, h=H, w=W) + x = out + + # output + x = x.flatten(2) + x = self.o(x) + return x + + +class MotionerAttentionBlock(nn.Module): + + def __init__(self, + dim, + ffn_dim, + num_heads, + window_size=(-1, -1), + qk_norm=True, + cross_attn_norm=False, + eps=1e-6, + self_attn_block="SelfAttention"): + super().__init__() + self.dim = dim + self.ffn_dim = ffn_dim + self.num_heads = num_heads + self.window_size = window_size + self.qk_norm = qk_norm + self.cross_attn_norm = cross_attn_norm + self.eps = eps + + # layers + self.norm1 = LayerNorm(dim, eps) + if self_attn_block == "SelfAttention": + self.self_attn = SelfAttention(dim, num_heads, window_size, qk_norm, + eps) + elif self_attn_block == "SwinSelfAttention": + self.self_attn = SwinSelfAttention(dim, num_heads, window_size, + qk_norm, eps) + elif self_attn_block == "CasualSelfAttention": + self.self_attn = CasualSelfAttention(dim, num_heads, window_size, + qk_norm, eps) + + self.norm2 = LayerNorm(dim, eps) + self.ffn = nn.Sequential( + nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'), + nn.Linear(ffn_dim, dim)) + + def forward( + self, + x, + seq_lens, + grid_sizes, + freqs, + ): + # self-attention + y = self.self_attn(self.norm1(x).float(), seq_lens, grid_sizes, freqs) + x = x + y + y = self.ffn(self.norm2(x).float()) + x = x + y + return x + + +class Head(nn.Module): + + def __init__(self, dim, out_dim, patch_size, eps=1e-6): + super().__init__() + self.dim = dim + self.out_dim = out_dim + self.patch_size = patch_size + self.eps = eps + + # layers + out_dim = math.prod(patch_size) * out_dim + self.norm = LayerNorm(dim, eps) + self.head = nn.Linear(dim, out_dim) + + def forward(self, x): + x = self.head(self.norm(x)) + return x + + +class MotionerTransformers(nn.Module, PeftAdapterMixin): + + def __init__( + self, + patch_size=(1, 2, 2), + in_dim=16, + dim=2048, + ffn_dim=8192, + freq_dim=256, + out_dim=16, + num_heads=16, + num_layers=32, + window_size=(-1, -1), + qk_norm=True, + cross_attn_norm=False, + eps=1e-6, + self_attn_block="SelfAttention", + motion_token_num=1024, + enable_tsm=False, + motion_stride=4, + expand_ratio=2, + trainable_token_pos_emb=False, + ): + super().__init__() + self.patch_size = patch_size + self.in_dim = in_dim + self.dim = dim + self.ffn_dim = ffn_dim + self.freq_dim = freq_dim + self.out_dim = out_dim + self.num_heads = num_heads + self.num_layers = num_layers + self.window_size = window_size + self.qk_norm = qk_norm + self.cross_attn_norm = cross_attn_norm + self.eps = eps + + self.enable_tsm = enable_tsm + self.motion_stride = motion_stride + self.expand_ratio = expand_ratio + self.sample_c = self.patch_size[0] + + # embeddings + self.patch_embedding = nn.Conv3d( + in_dim, dim, kernel_size=patch_size, stride=patch_size) + + # blocks + self.blocks = nn.ModuleList([ + MotionerAttentionBlock( + dim, + ffn_dim, + num_heads, + window_size, + qk_norm, + cross_attn_norm, + eps, + self_attn_block=self_attn_block) for _ in range(num_layers) + ]) + + # buffers (don't use register_buffer otherwise dtype will be changed in to()) + assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0 + d = dim // num_heads + self.freqs = torch.cat([ + rope_params(1024, d - 4 * (d // 6)), + rope_params(1024, 2 * (d // 6)), + rope_params(1024, 2 * (d // 6)) + ], + dim=1) + + self.gradient_checkpointing = False + + self.motion_side_len = int(math.sqrt(motion_token_num)) + assert self.motion_side_len**2 == motion_token_num + self.token = nn.Parameter( + torch.zeros(1, motion_token_num, dim).contiguous()) + + self.trainable_token_pos_emb = trainable_token_pos_emb + if trainable_token_pos_emb: + x = torch.zeros([1, motion_token_num, num_heads, d]) + x[..., ::2] = 1 + + gride_sizes = [[ + torch.tensor([0, 0, 0]).unsqueeze(0).repeat(1, 1), + torch.tensor([1, self.motion_side_len, + self.motion_side_len]).unsqueeze(0).repeat(1, 1), + torch.tensor([1, self.motion_side_len, + self.motion_side_len]).unsqueeze(0).repeat(1, 1), + ]] + token_freqs = rope_apply(x, gride_sizes, self.freqs) + token_freqs = token_freqs[0, :, 0].reshape(motion_token_num, -1, 2) + token_freqs = token_freqs * 0.01 + self.token_freqs = torch.nn.Parameter(token_freqs) + + def after_patch_embedding(self, x): + return x + + def forward( + self, + x, + ): + """ + x: A list of videos each with shape [C, T, H, W]. + t: [B]. + context: A list of text embeddings each with shape [L, C]. + """ + # params + motion_frames = x[0].shape[1] + device = self.patch_embedding.weight.device + freqs = self.freqs + if freqs.device != device: + freqs = freqs.to(device) + + if self.trainable_token_pos_emb: + with amp.autocast(dtype=torch.float64): + token_freqs = self.token_freqs.to(torch.float64) + token_freqs = token_freqs / token_freqs.norm( + dim=-1, keepdim=True) + freqs = [freqs, torch.view_as_complex(token_freqs)] + + if self.enable_tsm: + sample_idx = [ + sample_indices( + u.shape[1], + stride=self.motion_stride, + expand_ratio=self.expand_ratio, + c=self.sample_c) for u in x + ] + x = [ + torch.flip(torch.flip(u, [1])[:, idx], [1]) + for idx, u in zip(sample_idx, x) + ] + + # embeddings + x = [self.patch_embedding(u.unsqueeze(0)) for u in x] + x = self.after_patch_embedding(x) + + seq_f, seq_h, seq_w = x[0].shape[-3:] + batch_size = len(x) + if not self.enable_tsm: + grid_sizes = torch.stack( + [torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) + grid_sizes = [[ + torch.zeros_like(grid_sizes), grid_sizes, grid_sizes + ]] + seq_f = 0 + else: + grid_sizes = [] + for idx in sample_idx[0][::-1][::self.sample_c]: + tsm_frame_grid_sizes = [[ + torch.tensor([idx, 0, + 0]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([idx + 1, seq_h, + seq_w]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([1, seq_h, + seq_w]).unsqueeze(0).repeat(batch_size, 1), + ]] + grid_sizes += tsm_frame_grid_sizes + seq_f = sample_idx[0][-1] + 1 + + x = [u.flatten(2).transpose(1, 2) for u in x] + seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long) + x = torch.cat([u for u in x]) + + batch_size = len(x) + + token_grid_sizes = [[ + torch.tensor([seq_f, 0, 0]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor( + [seq_f + 1, self.motion_side_len, + self.motion_side_len]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor( + [1 if not self.trainable_token_pos_emb else -1, seq_h, + seq_w]).unsqueeze(0).repeat(batch_size, 1), + ] # 第三行代表rope emb的想要覆盖到的范围 + ] + + grid_sizes = grid_sizes + token_grid_sizes + token_unpatch_grid_sizes = torch.stack([ + torch.tensor([1, 32, 32], dtype=torch.long) + for b in range(batch_size) + ]) + token_len = self.token.shape[1] + token = self.token.clone().repeat(x.shape[0], 1, 1).contiguous() + seq_lens = seq_lens + torch.tensor([t.size(0) for t in token], + dtype=torch.long) + x = torch.cat([x, token], dim=1) + # arguments + kwargs = dict( + seq_lens=seq_lens, + grid_sizes=grid_sizes, + freqs=freqs, + ) + + # grad ckpt args + def create_custom_forward(module, return_dict=None): + + def custom_forward(*inputs, **kwargs): + if return_dict is not None: + return module(*inputs, **kwargs, return_dict=return_dict) + else: + return module(*inputs, **kwargs) + + return custom_forward + + ckpt_kwargs: Dict[str, Any] = ({ + "use_reentrant": False + } if is_torch_version(">=", "1.11.0") else {}) + + for idx, block in enumerate(self.blocks): + if self.training and self.gradient_checkpointing: + x = torch.utils.checkpoint.checkpoint( + create_custom_forward(block), + x, + **kwargs, + **ckpt_kwargs, + ) + else: + x = block(x, **kwargs) + # head + out = x[:, -token_len:] + return out + + def unpatchify(self, x, grid_sizes): + c = self.out_dim + out = [] + for u, v in zip(x, grid_sizes.tolist()): + u = u[:math.prod(v)].view(*v, *self.patch_size, c) + u = torch.einsum('fhwpqrc->cfphqwr', u) + u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)]) + out.append(u) + return out + + def init_weights(self): + # basic init + for m in self.modules(): + if isinstance(m, nn.Linear): + nn.init.xavier_uniform_(m.weight) + if m.bias is not None: + nn.init.zeros_(m.bias) + + # init embeddings + nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1)) + + +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 + *args, + **kwargs): + super().__init__(*args, **kwargs) + 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 = torch.tensor( + zip_frame_buckets, dtype=torch.long) + + self.inner_dim = inner_dim + self.num_heads = num_heads + + assert (inner_dim % + num_heads) == 0 and (inner_dim // num_heads) % 2 == 0 + d = inner_dim // num_heads + self.freqs = torch.cat([ + rope_params(1024, d - 4 * (d // 6)), + rope_params(1024, 2 * (d // 6)), + rope_params(1024, 2 * (d // 6)) + ], + dim=1) + self.drop_mode = drop_mode + + def forward(self, motion_latents, add_last_motion=2): + motion_frames = motion_latents[0].shape[1] + mot = [] + mot_remb = [] + for m in motion_latents: + lat_height, lat_width = m.shape[2], m.shape[3] + padd_lat = torch.zeros(16, self.zip_frame_buckets.sum(), lat_height, + lat_width).to( + device=m.device, dtype=m.dtype) + overlap_frame = min(padd_lat.shape[1], m.shape[1]) + if overlap_frame > 0: + padd_lat[:, -overlap_frame:] = m[:, -overlap_frame:] + + if add_last_motion < 2 and self.drop_mode != "drop": + zero_end_frame = self.zip_frame_buckets[:self.zip_frame_buckets. + __len__() - + add_last_motion - + 1].sum() + padd_lat[:, -zero_end_frame:] = 0 + + padd_lat = padd_lat.unsqueeze(0) + clean_latents_4x, clean_latents_2x, clean_latents_post = padd_lat[:, :, -self.zip_frame_buckets.sum( + ):, :, :].split( + list(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).flatten( + 2).transpose(1, 2) + clean_latents_4x = self.proj_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 + start_time_id = -(self.zip_frame_buckets[:1].sum()) + end_time_id = start_time_id + self.zip_frame_buckets[0] + grid_sizes = [] if add_last_motion < 2 and self.drop_mode == "drop" else \ + [ + [torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), + torch.tensor([end_time_id, lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), + torch.tensor([self.zip_frame_buckets[0], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), ] + ] + + start_time_id = -(self.zip_frame_buckets[:2].sum()) + end_time_id = start_time_id + self.zip_frame_buckets[1] // 2 + grid_sizes_2x = [] if add_last_motion < 1 and self.drop_mode == "drop" else \ + [ + [torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), + torch.tensor([end_time_id, lat_height // 4, lat_width // 4]).unsqueeze(0).repeat(1, 1), + torch.tensor([self.zip_frame_buckets[1], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), ] + ] + + start_time_id = -(self.zip_frame_buckets[:3].sum()) + end_time_id = start_time_id + self.zip_frame_buckets[2] // 4 + grid_sizes_4x = [[ + torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), + torch.tensor([end_time_id, lat_height // 8, + lat_width // 8]).unsqueeze(0).repeat(1, 1), + torch.tensor([ + self.zip_frame_buckets[2], lat_height // 2, lat_width // 2 + ]).unsqueeze(0).repeat(1, 1), + ]] + + grid_sizes = grid_sizes + grid_sizes_2x + grid_sizes_4x + + motion_rope_emb = rope_precompute( + motion_lat.detach().view(1, motion_lat.shape[1], self.num_heads, + self.inner_dim // self.num_heads), + grid_sizes, + self.freqs, + start=None) + + mot.append(motion_lat) + mot_remb.append(motion_rope_emb) + return mot, mot_remb + + +def sample_indices(N, stride, expand_ratio, c): + indices = [] + current_start = 0 + + while current_start < N: + bucket_width = int(stride * (expand_ratio**(len(indices) / stride))) + + interval = int(bucket_width / stride * c) + current_end = min(N, current_start + bucket_width) + bucket_samples = [] + for i in range(current_end - 1, current_start - 1, -interval): + for near in range(c): + bucket_samples.append(i - near) + + indices += bucket_samples[::-1] + current_start += bucket_width + + return indices + + +if __name__ == '__main__': + device = "cuda" + model = FramePackMotioner(inner_dim=1024) + batch_size = 2 + num_frame, height, width = (28, 32, 32) + single_input = torch.ones([16, num_frame, height, width], device=device) + for i in range(num_frame): + single_input[:, num_frame - 1 - i] *= i + x = [single_input] * batch_size + model.forward(x) diff --git a/wanvideo/modules/s2v/s2v_utils.py b/wanvideo/modules/s2v/s2v_utils.py new file mode 100644 index 0000000..68644a2 --- /dev/null +++ b/wanvideo/modules/s2v/s2v_utils.py @@ -0,0 +1,70 @@ +# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. +import numpy as np +import torch + + +def rope_precompute(x, grid_sizes, freqs, start=None): + b, s, n, c = x.size(0), x.size(1), x.size(2), x.size(3) // 2 + + # split freqs + if type(freqs) is list: + trainable_freqs = freqs[1] + freqs = freqs[0] + freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1) + + # loop over samples + output = torch.view_as_complex(x.detach().reshape(b, s, n, -1, + 2).to(torch.float64)) + seq_bucket = [0] + if not type(grid_sizes) is list: + grid_sizes = [grid_sizes] + for g in grid_sizes: + if not type(g) is list: + g = [torch.zeros_like(g), g] + batch_size = g[0].shape[0] + for i in range(batch_size): + if start is None: + f_o, h_o, w_o = g[0][i] + else: + f_o, h_o, w_o = start[i] + + f, h, w = g[1][i] + t_f, t_h, t_w = g[2][i] + seq_f, seq_h, seq_w = f - f_o, h - h_o, w - w_o + seq_len = int(seq_f * seq_h * seq_w) + if seq_len > 0: + if t_f > 0: + factor_f, factor_h, factor_w = (t_f / seq_f).item(), ( + t_h / seq_h).item(), (t_w / seq_w).item() + # Generate a list of seq_f integers starting from f_o and ending at math.ceil(factor_f * seq_f.item() + f_o.item()) + if f_o >= 0: + f_sam = np.linspace(f_o.item(), (t_f + f_o).item() - 1, + seq_f).astype(int).tolist() + else: + f_sam = np.linspace(-f_o.item(), + (-t_f - f_o).item() + 1, + seq_f).astype(int).tolist() + h_sam = np.linspace(h_o.item(), (t_h + h_o).item() - 1, + seq_h).astype(int).tolist() + w_sam = np.linspace(w_o.item(), (t_w + w_o).item() - 1, + seq_w).astype(int).tolist() + + assert f_o * f >= 0 and h_o * h >= 0 and w_o * w >= 0 + freqs_0 = freqs[0][f_sam] if f_o >= 0 else freqs[0][ + f_sam].conj() + freqs_0 = freqs_0.view(seq_f, 1, 1, -1) + + freqs_i = torch.cat([ + freqs_0.expand(seq_f, seq_h, seq_w, -1), + freqs[1][h_sam].view(1, seq_h, 1, -1).expand( + seq_f, seq_h, seq_w, -1), + freqs[2][w_sam].view(1, 1, seq_w, -1).expand( + seq_f, seq_h, seq_w, -1), + ], + dim=-1).reshape(seq_len, 1, -1) + elif t_f < 0: + freqs_i = trainable_freqs.unsqueeze(1) + # apply rotary embedding + output[i, seq_bucket[-1]:seq_bucket[-1] + seq_len] = freqs_i + seq_bucket.append(seq_bucket[-1] + seq_len) + return output From 63d4b6aadaae543f96d101788122563f3e2ba0c8 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 26 Aug 2025 19:07:38 +0300 Subject: [PATCH 02/31] Update nodes.py --- s2v/nodes.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/s2v/nodes.py b/s2v/nodes.py index bb4464e..589cefb 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -52,9 +52,10 @@ class WanVideoAddAudioEmbeds: return {"required": { "embeds": ("WANVIDIMAGE_EMBEDS",), "audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",), - "input_fps": ("FLOAT", {"default": 30.0, "min": 1.0, "max": 120.0, "step": 1.0, "tooltip": "Frames per second for the audio"}), + "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"}), - "frames": ("INT", {"default": 81, "min": 1, "max": 120, "step": 1, "tooltip": "Number of frames to process"}) + "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"}) } } @@ -63,7 +64,7 @@ class WanVideoAddAudioEmbeds: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, input_fps, output_fps, frames, audio_encoder_output): + 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"] @@ -80,7 +81,7 @@ class WanVideoAddAudioEmbeds: audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps( audio_feat, - fps=output_fps, + fps=bucket_fps, batch_frames=frames ) From 3c79851230c9ab042f52c8e176f349ecc3e51a64 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 26 Aug 2025 21:30:10 +0300 Subject: [PATCH 03/31] ref latent --- nodes.py | 8 ++- nodes_model_loading.py | 3 +- s2v/nodes.py | 23 +++---- wanvideo/modules/model.py | 131 +++++++++++++++++++++++++++++--------- 4 files changed, 120 insertions(+), 45 deletions(-) diff --git a/nodes.py b/nodes.py index 1abe0a9..e6d0cc3 100644 --- a/nodes.py +++ b/nodes.py @@ -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 diff --git a/nodes_model_loading.py b/nodes_model_loading.py index cbabd33..31e6b69 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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", } diff --git a/s2v/nodes.py b/s2v/nodes.py index 589cefb..fcd67a4 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -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 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index d232e89..2b739e1 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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] From 96f7f6accddd3a080953d0744be12c6967d5b444 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 26 Aug 2025 22:02:02 +0300 Subject: [PATCH 04/31] fix normal model loading --- nodes.py | 4 +- s2v/nodes.py | 6 ++- wanvideo/modules/model.py | 98 +++++++++++++++++---------------------- 3 files changed, 50 insertions(+), 58 deletions(-) diff --git a/nodes.py b/nodes.py index e6d0cc3..62f6150 100644 --- a/nodes.py +++ b/nodes.py @@ -2218,6 +2218,7 @@ class WanVideoSampler: 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_audio_scale = s2v_audio_embeds["audio_scale"] s2v_ref_latent = s2v_audio_embeds["ref_latent"] if s2v_ref_latent is not None: s2v_ref_latent = s2v_ref_latent.to(device, dtype) @@ -2683,7 +2684,8 @@ class WanVideoSampler: "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 # official speech-to-video reference latent + "s2v_ref_latent": s2v_ref_latent, # speech-to-video reference latent + "s2v_audio_scale": s2v_audio_scale if s2v_audio_input is not None else 1.0 # speech-to-video audio scale } batch_size = 1 diff --git a/s2v/nodes.py b/s2v/nodes.py index fcd67a4..6c840b7 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -53,6 +53,7 @@ class WanVideoAddAudioEmbeds: "embeds": ("WANVIDIMAGE_EMBEDS",), "audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",), "frames": ("INT", {"default": 81, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames to process"}), + "audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for audio embeddings"}) }, "optional": { "ref_latent": ("LATENT",) @@ -65,7 +66,7 @@ class WanVideoAddAudioEmbeds: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, frames, audio_encoder_output, ref_latent=None): + def add(self, embeds, frames, audio_encoder_output, audio_scale, 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] @@ -98,7 +99,8 @@ class WanVideoAddAudioEmbeds: new_entry = { "audio_embed_bucket": audio_embed_bucket, "num_repeat": num_repeat, - "ref_latent": ref_latent["samples"] if ref_latent is not None else None + "ref_latent": ref_latent["samples"] if ref_latent is not None else None, + "audio_scale": audio_scale } updated = dict(embeds) updated["audio_embeds"] = new_entry diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 2b739e1..a77172b 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1591,35 +1591,37 @@ class WanModel(torch.nn.Module): self.block_mask=None #S2V + self.zero_timestep = None if cond_dim > 0: self.cond_encoder = nn.Conv3d( cond_dim, self.dim, kernel_size=self.patch_size, stride=self.patch_size) - 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", - ) - self.adain_mode = adain_mode - self.zero_timestep = zero_timestep + 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", + ) + self.adain_mode = adain_mode + self.zero_timestep = zero_timestep - self.trainable_cond_mask = nn.Embedding(3, self.dim) + self.trainable_cond_mask = nn.Embedding(3, self.dim) @staticmethod def _prepare_blockwise_causal_attn_mask( @@ -1772,45 +1774,30 @@ class WanModel(torch.nn.Module): return hints - def audio_injector_forward(self, block_idx, hidden_states, merged_audio_emb): + 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] - audio_emb = merged_audio_emb # b f n c - num_frames = audio_emb.shape[1] + num_frames = audio_emb.shape[1]# b f n c - input_hidden_states = hidden_states[:, :self.original_seq_len].clone() # b (f h w) c - input_hidden_states = rearrange( - input_hidden_states, "b (t n) c -> (b t) n c", t=num_frames) + 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") - adain_hidden_states = self.audio_injector.injector_adain_layers[ - audio_attn_id]( - input_hidden_states, temb=audio_emb_global[:, 0]) - attn_hidden_states = adain_hidden_states + 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_hidden_states = self.audio_injector.injector_pre_norm_feat[ - audio_attn_id]( - input_hidden_states) - audio_emb = rearrange( - audio_emb, "b t n c -> (b t) n c", t=num_frames) - attn_audio_emb = audio_emb - residual_out = self.audio_injector.injector[audio_attn_id]( - x=attn_hidden_states, - context=attn_audio_emb, - context_lens=torch.ones( - attn_hidden_states.shape[0], - dtype=torch.long, - device=attn_hidden_states.device) * attn_audio_emb.shape[1]) - residual_out = rearrange( - residual_out, "(b t) n c -> b (t n) c", t=num_frames) - hidden_states[:, :self. - original_seq_len] = hidden_states[:, :self. - original_seq_len] + residual_out + attn_x = self.audio_injector.injector_pre_norm_feat[audio_attn_id](input_x) - return hidden_states + 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 forward( self, @@ -1857,7 +1844,8 @@ class WanModel(torch.nn.Module): mtv_freqs=None, mtv_strength=1.0, s2v_audio_input=None, - s2v_ref_latent=None + s2v_ref_latent=None, + s2v_audio_scale=1.0 ): r""" @@ -2451,7 +2439,7 @@ class WanModel(torch.nn.Module): 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) #s2v + 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 From 8e65fae3c5efc99d32cba8cf58cacde43c73c91c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 26 Aug 2025 22:02:36 +0300 Subject: [PATCH 05/31] Update model.py --- wanvideo/modules/model.py | 1 + 1 file changed, 1 insertion(+) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index a77172b..40b25bb 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -787,6 +787,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): From e85e1ff9e6860a0a0dbb4f5a14c81567c94acc1b Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 26 Aug 2025 22:05:22 +0300 Subject: [PATCH 06/31] Update model.py --- wanvideo/modules/model.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 40b25bb..88b35c9 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1592,7 +1592,7 @@ class WanModel(torch.nn.Module): self.block_mask=None #S2V - self.zero_timestep = None + self.zero_timestep = self.audio_injector = None if cond_dim > 0: self.cond_encoder = nn.Conv3d( cond_dim, @@ -1619,10 +1619,9 @@ class WanModel(torch.nn.Module): adain_dim=self.dim, 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) + self.adain_mode = adain_mode + self.zero_timestep = zero_timestep @staticmethod def _prepare_blockwise_causal_attn_mask( From 5d17484cc52d689409bcec7979abd5386e3ecdb1 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 02:25:39 +0300 Subject: [PATCH 07/31] continue --- nodes.py | 11 +- s2v/nodes.py | 2 - wanvideo/modules/model.py | 234 +++++++++++++++++++++++++++++-- wanvideo/modules/s2v/motioner.py | 136 ++++-------------- 4 files changed, 265 insertions(+), 118 deletions(-) diff --git a/nodes.py b/nodes.py index 359fbe2..dd70eb8 100644 --- a/nodes.py +++ b/nodes.py @@ -2226,6 +2226,7 @@ class WanVideoSampler: 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 = s2v_audio_input[..., 0:image_embeds["num_frames"]] #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) @@ -2509,7 +2510,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): nonlocal transformer z = z.to(dtype) autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear) @@ -3202,6 +3203,10 @@ 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: + partial_s2v_audio_input = s2v_audio_input[..., c] + partial_add_cond = None if add_cond is not None: partial_add_cond = add_cond[:, :, c].to(device, dtype) @@ -3219,7 +3224,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) if cache_args is not None: self.window_tracker.cache_states[window_id] = new_teacache @@ -3653,7 +3658,7 @@ class WanVideoSampler: 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, diff --git a/s2v/nodes.py b/s2v/nodes.py index 6c840b7..9bd8ab1 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -92,8 +92,6 @@ class WanVideoAddAudioEmbeds: elif len(audio_embed_bucket.shape) == 4: audio_embed_bucket = audio_embed_bucket.permute(0, 2, 3, 1) - audio_embed_bucket = audio_embed_bucket[..., 0:frames] - print("audio_embed_bucket", audio_embed_bucket.shape) new_entry = { diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 88b35c9..03b1362 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -30,6 +30,8 @@ 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 diffusers.models.attention import AdaLayerNorm __all__ = ['WanModel'] @@ -191,7 +193,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) @@ -1302,7 +1303,6 @@ class AudioInjector_WAN(nn.Module): adain_dim=2048, need_adain_ont=False): super().__init__() - num_injector_layers = len(inject_layer) self.injected_block_id = {} audio_injector_id = 0 for mod_name, mod in zip(all_modules_names, all_modules): @@ -1623,6 +1623,48 @@ class WanModel(torch.nn.Module): self.adain_mode = adain_mode self.zero_timestep = zero_timestep + # init motioner + enable_framepack = False + enable_motioner = False + add_last_motion = False + if enable_motioner and enable_framepack: + raise ValueError( + "enable_motioner and enable_framepack are mutually exclusive, please set one of them to False" + ) + self.enable_motioner = enable_motioner + self.add_last_motion = add_last_motion + if enable_motioner: + motioner_dim = 2048 + self.motioner = MotionerTransformers( + patch_size=(2, 4, 4), + dim=motioner_dim, + ffn_dim=motioner_dim, + freq_dim=256, + out_dim=16, + num_heads=16, + num_layers=13, + window_size=(-1, -1), + qk_norm=True, + cross_attn_norm=False, + eps=1e-6, + motion_token_num=4, + enable_tsm=False, + motion_stride=4, + expand_ratio=2, + trainable_token_pos_emb=False, + ) + self.zip_motion_out = torch.nn.Sequential( + WanLayerNorm(motioner_dim), + zero_module(nn.Linear(motioner_dim, self.dim))) + + self.enable_framepack = enable_framepack + if enable_framepack: + self.frame_packer = FramePackMotioner( + inner_dim=self.dim, + num_heads=self.num_heads, + zip_frame_buckets=[1, 2, 16], + drop_mode='padd') + @staticmethod def _prepare_blockwise_causal_attn_mask( device: torch.device | str, num_frames: int = 21, @@ -1773,6 +1815,151 @@ class WanModel(torch.nn.Module): block.to(self.offload_device, non_blocking=self.use_non_blocking) return hints + + def process_motion(self, motion_latents, drop_motion_frames=False): + if drop_motion_frames or motion_latents[0].shape[1] == 0: + return [], [] + self.lat_motion_frames = motion_latents[0].shape[1] + mot = [self.patch_embedding(m.unsqueeze(0)) for m in motion_latents] + batch_size = len(mot) + + mot_remb = [] + flattern_mot = [] + for bs in range(batch_size): + height, width = mot[bs].shape[3], mot[bs].shape[4] + flat_mot = mot[bs].flatten(2).transpose(1, 2).contiguous() + motion_grid_sizes = [[ + torch.tensor([-self.lat_motion_frames, 0, + 0]).unsqueeze(0).repeat(1, 1), + torch.tensor([0, height, width]).unsqueeze(0).repeat(1, 1), + torch.tensor([self.lat_motion_frames, height, + width]).unsqueeze(0).repeat(1, 1) + ]] + motion_rope_emb = rope_precompute( + flat_mot.detach().view(1, flat_mot.shape[1], self.num_heads, + self.dim // self.num_heads), + motion_grid_sizes, + self.freqs, + start=None) + mot_remb.append(motion_rope_emb) + flattern_mot.append(flat_mot) + return flattern_mot, mot_remb + + def process_motion_frame_pack(self, + motion_latents, + drop_motion_frames=False, + add_last_motion=2): + flattern_mot, mot_remb = self.frame_packer(motion_latents, + add_last_motion) + if drop_motion_frames: + return [m[:, :0] for m in flattern_mot + ], [m[:, :0] for m in mot_remb] + else: + return flattern_mot, mot_remb + + def process_motion_transformer_motioner(self, + motion_latents, + drop_motion_frames=False, + add_last_motion=True): + batch_size, height, width = len( + motion_latents), motion_latents[0].shape[2] // self.patch_size[ + 1], motion_latents[0].shape[3] // self.patch_size[2] + + freqs = self.freqs + device = self.patch_embedding.weight.device + if freqs.device != device: + freqs = freqs.to(device) + if self.trainable_token_pos_emb: + token_freqs = self.token_freqs.to(torch.float64) + token_freqs = token_freqs / token_freqs.norm( + dim=-1, keepdim=True) + freqs = [freqs, torch.view_as_complex(token_freqs)] + + if not drop_motion_frames and add_last_motion: + last_motion_latent = [u[:, -1:] for u in motion_latents] + last_mot = [ + self.patch_embedding(m.unsqueeze(0)) for m in last_motion_latent + ] + last_mot = [m.flatten(2).transpose(1, 2) for m in last_mot] + last_mot = torch.cat(last_mot) + gride_sizes = [[ + torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([0, height, + width]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([1, height, + width]).unsqueeze(0).repeat(batch_size, 1) + ]] + else: + last_mot = torch.zeros([batch_size, 0, self.dim], + device=motion_latents[0].device, + dtype=motion_latents[0].dtype) + gride_sizes = [] + + zip_motion = self.motioner(motion_latents) + zip_motion = self.zip_motion_out(zip_motion) + if drop_motion_frames: + zip_motion = zip_motion * 0.0 + zip_motion_grid_sizes = [[ + torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([ + 0, self.motioner.motion_side_len, self.motioner.motion_side_len + ]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor( + [1 if not self.trainable_token_pos_emb else -1, height, + width]).unsqueeze(0).repeat(batch_size, 1), + ]] + + mot = torch.cat([last_mot, zip_motion], dim=1) + gride_sizes = gride_sizes + zip_motion_grid_sizes + + motion_rope_emb = rope_precompute( + mot.detach().view(batch_size, mot.shape[1], self.num_heads, + self.dim // self.num_heads), + gride_sizes, + freqs, + start=None) + return [m.unsqueeze(0) for m in mot + ], [r.unsqueeze(0) for r in motion_rope_emb] + + def inject_motion(self, + x, + seq_lens, + rope_embs, + mask_input, + motion_latents, + drop_motion_frames=False, + add_last_motion=True): + # inject the motion frames token to the hidden states + if self.enable_motioner: + mot, mot_remb = self.process_motion_transformer_motioner( + motion_latents, + drop_motion_frames=drop_motion_frames, + add_last_motion=add_last_motion) + elif self.enable_framepack: + mot, mot_remb = self.process_motion_frame_pack( + motion_latents, + drop_motion_frames=drop_motion_frames, + add_last_motion=add_last_motion) + else: + mot, mot_remb = self.process_motion( + motion_latents, drop_motion_frames=drop_motion_frames) + + if len(mot) > 0: + x = [torch.cat([u, m], dim=1) for u, m in zip(x, mot)] + seq_lens = seq_lens + torch.tensor([r.size(1) for r in mot], + dtype=torch.long) + rope_embs = [ + torch.cat([u, m], dim=1) for u, m in zip(rope_embs, mot_remb) + ] + mask_input = [ + torch.cat([ + m, 2 * torch.ones([1, u.shape[1] - m.shape[1]], + device=m.device, + dtype=m.dtype) + ], + dim=1) for m, u in zip(mask_input, x) + ] + return x, seq_lens, rope_embs, mask_input def audio_injector_forward(self, block_idx, x, audio_emb, scale=1.0): if block_idx in self.audio_injector.injected_block_id.keys(): @@ -1845,7 +2032,8 @@ class WanModel(torch.nn.Module): mtv_strength=1.0, s2v_audio_input=None, s2v_ref_latent=None, - s2v_audio_scale=1.0 + s2v_audio_scale=1.0, + motion_latents=None ): r""" @@ -1893,8 +2081,8 @@ class WanModel(torch.nn.Module): #s2v if self.model_type == 's2v' and s2v_audio_input is not None: - motion_frames=[17, 5] - + #motion_frames=[17, 5] + motion_frames=[1, 0] s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, motion_frames[0]), s2v_audio_input], dim=-1) audio_emb_res = self.casual_audio_encoder(s2v_audio_input) @@ -2108,10 +2296,40 @@ class WanModel(torch.nn.Module): e0 = torch.cat([ e0.unsqueeze(2), zero_e0.unsqueeze(2).repeat(e0.size(0), 1, 1, 1) - ], - dim=2) + ], dim=2) e0 = [e0, self.original_seq_len] + mask_input = torch.zeros([1, x.shape[1]], dtype=torch.int32, device=x.device) + mask_input[:, self.original_seq_len:] = 1 + + + if motion_latents is not None: + # compute the rope embeddings for the input + #x = torch.cat(x) + b, s, n, d = x.size(0), x.size( + 1), self.num_heads, self.dim // self.num_heads + self.pre_compute_freqs = rope_precompute( + x.detach().view(b, s, n, d), grid_sizes, freqs, start=None) + + x = [u.unsqueeze(0) for u in x] + self.pre_compute_freqs = [ + u.unsqueeze(0) for u in self.pre_compute_freqs + ] + x, seq_lens, self.pre_compute_freqs, mask_input = self.inject_motion( + x, + seq_lens, + self.pre_compute_freqs, + mask_input, + motion_latents, + drop_motion_frames=self.drop_motion_frames, + add_last_motion=True) + + x = torch.cat(x, dim=0) + self.pre_compute_freqs = torch.cat(self.pre_compute_freqs, dim=0) + mask_input = torch.cat(mask_input, dim=0) + + x = x + self.trainable_cond_mask(mask_input).to(x.dtype) + 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 ) @@ -2503,7 +2721,7 @@ class WanModel(torch.nn.Module): 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 = 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] diff --git a/wanvideo/modules/s2v/motioner.py b/wanvideo/modules/s2v/motioner.py index 699c570..043e476 100644 --- a/wanvideo/modules/s2v/motioner.py +++ b/wanvideo/modules/s2v/motioner.py @@ -1,16 +1,16 @@ # Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. import math -from typing import Any, Dict, List, Literal, Optional, Union +from typing import Any, Dict import numpy as np import torch import torch.cuda.amp as amp import torch.nn as nn -from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin -from diffusers.utils import BaseOutput, is_torch_version +from diffusers.loaders import PeftAdapterMixin +from diffusers.utils import BaseOutput from einops import rearrange, repeat -from ..model import flash_attention +from ..attention import attention from .s2v_utils import rope_precompute @@ -172,12 +172,11 @@ class SelfAttention(nn.Module): q, k, v = qkv_fn(x) - x = flash_attention( + x = attention( q=rope_apply(q, grid_sizes, freqs), k=rope_apply(k, grid_sizes, freqs), v=v, - k_lens=seq_lens, - window_size=self.window_size) + k_lens=seq_lens) # output x = x.flatten(2) @@ -224,12 +223,7 @@ class SwinSelfAttention(SelfAttention): # q: b (t h w) n d # k: b (t h w) n d - out = flash_attention( - q=q, - k=k, - v=v, - # k_lens=torch.tensor([k.shape[1]] * k.shape[0], device=x.device, dtype=torch.long), - window_size=self.window_size) + out = attention(q, k, v) out = torch.cat([out, ref_v[:1]], axis=0) out = rearrange(out, '(b t) (h w) n d -> b (t h w) n d', t=T, h=H, w=W) x = out @@ -308,7 +302,7 @@ class CasualSelfAttention(SelfAttention): # k: b (t h w) n d outs = [] for i in range(q.shape[0]): - out = flash_attention( + out = attention( q=q[i:i + 1], k=k[i:i + 1], v=v[i:i + 1], @@ -479,10 +473,8 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin): gride_sizes = [[ torch.tensor([0, 0, 0]).unsqueeze(0).repeat(1, 1), - torch.tensor([1, self.motion_side_len, - self.motion_side_len]).unsqueeze(0).repeat(1, 1), - torch.tensor([1, self.motion_side_len, - self.motion_side_len]).unsqueeze(0).repeat(1, 1), + torch.tensor([1, self.motion_side_len, self.motion_side_len]).unsqueeze(0).repeat(1, 1), + torch.tensor([1, self.motion_side_len, self.motion_side_len]).unsqueeze(0).repeat(1, 1), ]] token_freqs = rope_apply(x, gride_sizes, self.freqs) token_freqs = token_freqs[0, :, 0].reshape(motion_token_num, -1, 2) @@ -502,7 +494,6 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin): context: A list of text embeddings each with shape [L, C]. """ # params - motion_frames = x[0].shape[1] device = self.patch_embedding.weight.device freqs = self.freqs if freqs.device != device: @@ -535,22 +526,16 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin): seq_f, seq_h, seq_w = x[0].shape[-3:] batch_size = len(x) if not self.enable_tsm: - grid_sizes = torch.stack( - [torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) - grid_sizes = [[ - torch.zeros_like(grid_sizes), grid_sizes, grid_sizes - ]] + grid_sizes = torch.stack([torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) + grid_sizes = [[torch.zeros_like(grid_sizes), grid_sizes, grid_sizes]] seq_f = 0 else: grid_sizes = [] for idx in sample_idx[0][::-1][::self.sample_c]: tsm_frame_grid_sizes = [[ - torch.tensor([idx, 0, - 0]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([idx + 1, seq_h, - seq_w]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([1, seq_h, - seq_w]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([idx, 0, 0]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([idx + 1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1), ]] grid_sizes += tsm_frame_grid_sizes seq_f = sample_idx[0][-1] + 1 @@ -563,14 +548,10 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin): token_grid_sizes = [[ torch.tensor([seq_f, 0, 0]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor( - [seq_f + 1, self.motion_side_len, - self.motion_side_len]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor( - [1 if not self.trainable_token_pos_emb else -1, seq_h, - seq_w]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([seq_f + 1, self.motion_side_len,self.motion_side_len]).unsqueeze(0).repeat(batch_size, 1), + torch.tensor([1 if not self.trainable_token_pos_emb else -1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1), ] # 第三行代表rope emb的想要覆盖到的范围 - ] + ] grid_sizes = grid_sizes + token_grid_sizes token_unpatch_grid_sizes = torch.stack([ @@ -589,31 +570,8 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin): freqs=freqs, ) - # grad ckpt args - def create_custom_forward(module, return_dict=None): - - def custom_forward(*inputs, **kwargs): - if return_dict is not None: - return module(*inputs, **kwargs, return_dict=return_dict) - else: - return module(*inputs, **kwargs) - - return custom_forward - - ckpt_kwargs: Dict[str, Any] = ({ - "use_reentrant": False - } if is_torch_version(">=", "1.11.0") else {}) - for idx, block in enumerate(self.blocks): - if self.training and self.gradient_checkpointing: - x = torch.utils.checkpoint.checkpoint( - create_custom_forward(block), - x, - **kwargs, - **ckpt_kwargs, - ) - else: - x = block(x, **kwargs) + x = block(x, **kwargs) # head out = x[:, -token_len:] return out @@ -628,18 +586,6 @@ class MotionerTransformers(nn.Module, PeftAdapterMixin): out.append(u) return out - def init_weights(self): - # basic init - for m in self.modules(): - if isinstance(m, nn.Linear): - nn.init.xavier_uniform_(m.weight) - if m.bias is not None: - nn.init.zeros_(m.bias) - - # init embeddings - nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1)) - - class FramePackMotioner(nn.Module): def __init__( @@ -677,7 +623,6 @@ class FramePackMotioner(nn.Module): self.drop_mode = drop_mode def forward(self, motion_latents, add_last_motion=2): - motion_frames = motion_latents[0].shape[1] mot = [] mot_remb = [] for m in motion_latents: @@ -702,21 +647,15 @@ class FramePackMotioner(nn.Module): list(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).flatten( - 2).transpose(1, 2) - clean_latents_4x = self.proj_4x(clean_latents_4x).flatten( - 2).transpose(1, 2) + clean_latents_post = self.proj(clean_latents_post).flatten(2).transpose(1, 2) + clean_latents_2x = self.proj_2x(clean_latents_2x).flatten(2).transpose(1, 2) + clean_latents_4x = self.proj_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 + 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) + motion_lat = torch.cat([clean_latents_post, clean_latents_2x, clean_latents_4x], dim=1) # rope start_time_id = -(self.zip_frame_buckets[:1].sum()) @@ -731,20 +670,18 @@ class FramePackMotioner(nn.Module): start_time_id = -(self.zip_frame_buckets[:2].sum()) end_time_id = start_time_id + self.zip_frame_buckets[1] // 2 grid_sizes_2x = [] if add_last_motion < 1 and self.drop_mode == "drop" else \ - [ - [torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), + [[ + torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), torch.tensor([end_time_id, lat_height // 4, lat_width // 4]).unsqueeze(0).repeat(1, 1), - torch.tensor([self.zip_frame_buckets[1], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), ] - ] + torch.tensor([self.zip_frame_buckets[1], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), + ]] start_time_id = -(self.zip_frame_buckets[:3].sum()) end_time_id = start_time_id + self.zip_frame_buckets[2] // 4 grid_sizes_4x = [[ torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), - torch.tensor([end_time_id, lat_height // 8, - lat_width // 8]).unsqueeze(0).repeat(1, 1), - torch.tensor([ - self.zip_frame_buckets[2], lat_height // 2, lat_width // 2 + torch.tensor([end_time_id, lat_height // 8, lat_width // 8]).unsqueeze(0).repeat(1, 1), + torch.tensor([self.zip_frame_buckets[2], lat_height // 2, lat_width // 2 ]).unsqueeze(0).repeat(1, 1), ]] @@ -781,14 +718,3 @@ def sample_indices(N, stride, expand_ratio, c): return indices - -if __name__ == '__main__': - device = "cuda" - model = FramePackMotioner(inner_dim=1024) - batch_size = 2 - num_frame, height, width = (28, 32, 32) - single_input = torch.ones([16, num_frame, height, width], device=device) - for i in range(num_frame): - single_input[:, num_frame - 1 - i] *= i - x = [single_input] * batch_size - model.forward(x) From a84d142cb2b28782491385609447d8ab848fd5fa Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 02:40:27 +0300 Subject: [PATCH 08/31] Update nodes.py --- nodes.py | 10 +++++++++- 1 file changed, 9 insertions(+), 1 deletion(-) diff --git a/nodes.py b/nodes.py index dd70eb8..81b74d4 100644 --- a/nodes.py +++ b/nodes.py @@ -3205,7 +3205,15 @@ class WanVideoSampler: partial_s2v_audio_input = None if s2v_audio_input is not None: - partial_s2v_audio_input = s2v_audio_input[..., c] + audio_indices = [] + max_audio_index = s2v_audio_input.shape[-1] + for i in c: + for j in range(4): + i_ = i * 4 + j + if i_ < max_audio_index: + audio_indices.append(i_) + print("audio_indices:", audio_indices) + partial_s2v_audio_input = s2v_audio_input[..., audio_indices] partial_add_cond = None if add_cond is not None: From 0538f5d6000d41072ad38cbc50be790918d5b881 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 02:41:24 +0300 Subject: [PATCH 09/31] Update nodes.py --- nodes.py | 1 - 1 file changed, 1 deletion(-) diff --git a/nodes.py b/nodes.py index 81b74d4..15c21b0 100644 --- a/nodes.py +++ b/nodes.py @@ -3212,7 +3212,6 @@ class WanVideoSampler: i_ = i * 4 + j if i_ < max_audio_index: audio_indices.append(i_) - print("audio_indices:", audio_indices) partial_s2v_audio_input = s2v_audio_input[..., audio_indices] partial_add_cond = None From a507a5763d4651e4add50e3cd62622a972b25e77 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 13:03:24 +0300 Subject: [PATCH 10/31] update WanVideoScheduler -node --- nodes.py | 16 ++++++++++------ nodes_utility.py | 1 + wanvideo/schedulers/__init__.py | 2 +- 3 files changed, 12 insertions(+), 7 deletions(-) diff --git a/nodes.py b/nodes.py index 1d318e4..0a7dac2 100644 --- a/nodes.py +++ b/nodes.py @@ -1608,7 +1608,7 @@ class WanVideoScheduler: #WIP EXPERIMENTAL = True def process(self, scheduler, steps, start_step, end_step, shift, unique_id, sigmas=None): - sample_scheduler, timesteps = get_scheduler( + sample_scheduler, timesteps, start_idx, end_idx = get_scheduler( scheduler, steps, start_step, end_step, shift, @@ -1642,8 +1642,12 @@ class WanVideoScheduler: #WIP ax.tick_params(axis='x', colors='white') # X tick color ax.tick_params(axis='y', colors='white') # Y tick color # Add split point if end_step is defined - if end_step != -1 and 0 <= end_step < len(sigmas_np): - ax.axvline(end_step, color='red', linestyle='--', linewidth=2, label='end_step split') + if end_idx != -1 and 0 <= end_idx < len(sigmas_np): + ax.axvline(end_idx, color='red', linestyle='--', linewidth=2, label='end_step split') + # Add split point if start_step is defined + if start_idx > 0 and 0 <= start_idx < len(sigmas_np): + ax.axvline(start_idx, color='green', linestyle='--', linewidth=2, label='start_step split') + if (end_idx != -1 and 0 <= end_idx < len(sigmas_np)) or (start_idx > 0 and 0 <= start_idx < len(sigmas_np)): ax.legend() plt.tight_layout() plt.savefig(buf, format='png') @@ -1815,7 +1819,7 @@ class WanVideoSampler: sample_scheduler = scheduler["sample_scheduler"] timesteps = scheduler["timesteps"] elif scheduler != "multitalk": - sample_scheduler, timesteps = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) log.info(f"sigmas: {sample_scheduler.sigmas}") else: timesteps = torch.tensor([1000, 750, 500, 250], device=device) @@ -2886,7 +2890,7 @@ class WanVideoSampler: # FreeInit noise reinitialization (after first iteration) if freeinit_args is not None and iter_idx > 0: # restart scheduler for each iteration - sample_scheduler, timesteps = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) # Re-apply start_step and end_step logic to timesteps and sigmas if end_step != -1: @@ -3414,7 +3418,7 @@ class WanVideoSampler: timesteps = [torch.tensor([t], device=device) for t in timesteps] timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps] else: - sample_scheduler, timesteps = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) timesteps = [torch.tensor([float(t)], device=device) for t in timesteps] + [torch.tensor([0.], device=device)] # sample videos diff --git a/nodes_utility.py b/nodes_utility.py index 09293a8..9af2016 100644 --- a/nodes_utility.py +++ b/nodes_utility.py @@ -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 diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index 407e107..c778fe1 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -139,4 +139,4 @@ def get_scheduler(scheduler, steps, start_step, end_step, shift, device, transfo if hasattr(sample_scheduler, 'timesteps'): sample_scheduler.timesteps = timesteps - return sample_scheduler, timesteps \ No newline at end of file + return sample_scheduler, timesteps, start_idx, end_idx \ No newline at end of file From 90c3bbb6c2e4ff5e05305e765d007d5e58428ce4 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 18:47:31 +0300 Subject: [PATCH 11/31] better context window indices, cleanup --- nodes.py | 30 +- wanvideo/modules/model.py | 268 ++++++----- wanvideo/modules/s2v/motioner.py | 720 ------------------------------ wanvideo/modules/s2v/s2v_utils.py | 70 --- 4 files changed, 184 insertions(+), 904 deletions(-) delete mode 100644 wanvideo/modules/s2v/motioner.py delete mode 100644 wanvideo/modules/s2v/s2v_utils.py diff --git a/nodes.py b/nodes.py index 0a7dac2..dd2810b 100644 --- a/nodes.py +++ b/nodes.py @@ -2224,15 +2224,18 @@ class WanVideoSampler: #region S2V s2v_audio_input = s2v_ref_latent = 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") s2v_audio_input = s2v_audio_embeds["audio_embed_bucket"].to(device, dtype) s2v_audio_scale = s2v_audio_embeds["audio_scale"] - 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_ref_latent = s2v_audio_embeds["ref_latent"].to(device, dtype) if "ref_latent" in s2v_audio_embeds else None + s2v_ref_motion = s2v_audio_embeds["ref_motion"].to(device, dtype) if "ref_motion" in s2v_audio_embeds else None s2v_audio_input = s2v_audio_input[..., 0:image_embeds["num_frames"]] + s2v_num_repeat = s2v_audio_embeds.get("num_repeat", 1) + vae = image_embeds.get("vae", None) + framepack = False #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) @@ -2516,7 +2519,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, s2v_audio_input=None): + mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None): nonlocal transformer z = z.to(dtype) autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear) @@ -2696,6 +2699,7 @@ class WanVideoSampler: "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 } @@ -3211,14 +3215,12 @@ class WanVideoSampler: partial_s2v_audio_input = None if s2v_audio_input is not None: - audio_indices = [] - max_audio_index = s2v_audio_input.shape[-1] - for i in c: - for j in range(4): - i_ = i * 4 + j - if i_ < max_audio_index: - audio_indices.append(i_) - partial_s2v_audio_input = s2v_audio_input[..., audio_indices] + 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_add_cond = None if add_cond is not None: @@ -3245,7 +3247,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: @@ -3544,7 +3546,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) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 03b1362..21f2a91 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -30,7 +30,61 @@ 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 .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 + dtype=None, device=None): + super().__init__() + self.proj = nn.Conv3d(16, inner_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), dtype=dtype, device=device) + self.proj_2x = nn.Conv3d(16, inner_dim, kernel_size=(2, 4, 4), stride=(2, 4, 4), dtype=dtype, device=device) + self.proj_4x = nn.Conv3d(16, inner_dim, kernel_size=(4, 8, 8), stride=(4, 8, 8), dtype=dtype, device=device) + 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(1, lat_height, lat_width, t_start=-1, device=motion_latents.device, dtype=motion_latents.dtype) + rope_2x = rope_embedder.rope_encode(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(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 @@ -1301,7 +1355,8 @@ class AudioInjector_WAN(nn.Module): root_net=None, enable_adain=False, adain_dim=2048, - need_adain_ont=False): + need_adain_ont=False, + attention_mode='sdpa'): super().__init__() self.injected_block_id = {} audio_injector_id = 0 @@ -1318,6 +1373,7 @@ class AudioInjector_WAN(nn.Module): 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([ @@ -1592,7 +1648,7 @@ class WanModel(torch.nn.Module): self.block_mask=None #S2V - self.zero_timestep = self.audio_injector = None + self.zero_timestep = self.audio_injector = self.trainable_cond_mask =None if cond_dim > 0: self.cond_encoder = nn.Conv3d( cond_dim, @@ -1618,6 +1674,7 @@ class WanModel(torch.nn.Module): 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.adain_mode = adain_mode @@ -1633,29 +1690,29 @@ class WanModel(torch.nn.Module): ) self.enable_motioner = enable_motioner self.add_last_motion = add_last_motion - if enable_motioner: - motioner_dim = 2048 - self.motioner = MotionerTransformers( - patch_size=(2, 4, 4), - dim=motioner_dim, - ffn_dim=motioner_dim, - freq_dim=256, - out_dim=16, - num_heads=16, - num_layers=13, - window_size=(-1, -1), - qk_norm=True, - cross_attn_norm=False, - eps=1e-6, - motion_token_num=4, - enable_tsm=False, - motion_stride=4, - expand_ratio=2, - trainable_token_pos_emb=False, - ) - self.zip_motion_out = torch.nn.Sequential( - WanLayerNorm(motioner_dim), - zero_module(nn.Linear(motioner_dim, self.dim))) + # if enable_motioner: + # motioner_dim = 2048 + # self.motioner = MotionerTransformers( + # patch_size=(2, 4, 4), + # dim=motioner_dim, + # ffn_dim=motioner_dim, + # freq_dim=256, + # out_dim=16, + # num_heads=16, + # num_layers=13, + # window_size=(-1, -1), + # qk_norm=True, + # cross_attn_norm=False, + # eps=1e-6, + # motion_token_num=4, + # enable_tsm=False, + # motion_stride=4, + # expand_ratio=2, + # trainable_token_pos_emb=False, + # ) + # self.zip_motion_out = torch.nn.Sequential( + # WanLayerNorm(motioner_dim), + # zero_module(nn.Linear(motioner_dim, self.dim))) self.enable_framepack = enable_framepack if enable_framepack: @@ -1663,7 +1720,9 @@ class WanModel(torch.nn.Module): inner_dim=self.dim, num_heads=self.num_heads, zip_frame_buckets=[1, 2, 16], - drop_mode='padd') + drop_mode='padd', + device=self.main_device, + dtype=self.dtype) @staticmethod def _prepare_blockwise_causal_attn_mask( @@ -1985,6 +2044,52 @@ class WanModel(torch.nn.Module): 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, @@ -2033,7 +2138,7 @@ class WanModel(torch.nn.Module): s2v_audio_input=None, s2v_ref_latent=None, s2v_audio_scale=1.0, - motion_latents=None + s2v_ref_motion=None ): r""" @@ -2081,6 +2186,8 @@ class WanModel(torch.nn.Module): #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 #motion_frames=[17, 5] motion_frames=[1, 0] s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, motion_frames[0]), s2v_audio_input], dim=-1) @@ -2093,7 +2200,6 @@ class WanModel(torch.nn.Module): audio_emb = audio_emb_res merged_audio_emb = audio_emb[:, motion_frames[1]:, :] - # params device = self.patch_embedding.weight.device @@ -2147,16 +2253,16 @@ class WanModel(torch.nn.Module): 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] + 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].size(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) @@ -2177,13 +2283,16 @@ class WanModel(torch.nn.Module): 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)).flatten(2).transpose(1, 2).to(x[0].dtype) + 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 - 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)] grid_sizes = grid_sizes @@ -2192,6 +2301,9 @@ class WanModel(torch.nn.Module): 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 @@ -2208,10 +2320,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 @@ -2220,37 +2331,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=30, device=x.device, dtype=x.dtype) + freqs = torch.cat([freqs, freqs_ref], dim=1) self.cached_freqs = freqs self.cached_shape = current_shape @@ -2261,6 +2349,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) @@ -2275,6 +2365,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 @@ -2299,37 +2398,6 @@ class WanModel(torch.nn.Module): ], dim=2) e0 = [e0, self.original_seq_len] - mask_input = torch.zeros([1, x.shape[1]], dtype=torch.int32, device=x.device) - mask_input[:, self.original_seq_len:] = 1 - - - if motion_latents is not None: - # compute the rope embeddings for the input - #x = torch.cat(x) - b, s, n, d = x.size(0), x.size( - 1), self.num_heads, self.dim // self.num_heads - self.pre_compute_freqs = rope_precompute( - x.detach().view(b, s, n, d), grid_sizes, freqs, start=None) - - x = [u.unsqueeze(0) for u in x] - self.pre_compute_freqs = [ - u.unsqueeze(0) for u in self.pre_compute_freqs - ] - x, seq_lens, self.pre_compute_freqs, mask_input = self.inject_motion( - x, - seq_lens, - self.pre_compute_freqs, - mask_input, - motion_latents, - drop_motion_frames=self.drop_motion_frames, - add_last_motion=True) - - x = torch.cat(x, dim=0) - self.pre_compute_freqs = torch.cat(self.pre_compute_freqs, dim=0) - mask_input = torch.cat(mask_input, dim=0) - - x = x + self.trainable_cond_mask(mask_input).to(x.dtype) - 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 ) @@ -2718,7 +2786,7 @@ class WanModel(torch.nn.Module): 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] diff --git a/wanvideo/modules/s2v/motioner.py b/wanvideo/modules/s2v/motioner.py deleted file mode 100644 index 043e476..0000000 --- a/wanvideo/modules/s2v/motioner.py +++ /dev/null @@ -1,720 +0,0 @@ -# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. -import math -from typing import Any, Dict - -import numpy as np -import torch -import torch.cuda.amp as amp -import torch.nn as nn -from diffusers.loaders import PeftAdapterMixin -from diffusers.utils import BaseOutput -from einops import rearrange, repeat - -from ..attention import attention -from .s2v_utils import rope_precompute - - -def sinusoidal_embedding_1d(dim, position): - # preprocess - assert dim % 2 == 0 - half = dim // 2 - position = position.type(torch.float64) - - # calculation - sinusoid = torch.outer( - position, torch.pow(10000, -torch.arange(half).to(position).div(half))) - x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1) - return x - - -@amp.autocast(enabled=False) -def rope_params(max_seq_len, dim, theta=10000): - assert dim % 2 == 0 - freqs = torch.outer( - torch.arange(max_seq_len), - 1.0 / torch.pow(theta, - torch.arange(0, dim, 2).to(torch.float64).div(dim))) - freqs = torch.polar(torch.ones_like(freqs), freqs) - return freqs - - -@amp.autocast(enabled=False) -def rope_apply(x, grid_sizes, freqs, start=None): - n, c = x.size(2), x.size(3) // 2 - - # split freqs - if type(freqs) is list: - trainable_freqs = freqs[1] - freqs = freqs[0] - freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1) - - # loop over samples - output = [] - output = x.clone() - seq_bucket = [0] - if not type(grid_sizes) is list: - grid_sizes = [grid_sizes] - for g in grid_sizes: - if not type(g) is list: - g = [torch.zeros_like(g), g] - batch_size = g[0].shape[0] - for i in range(batch_size): - if start is None: - f_o, h_o, w_o = g[0][i] - else: - f_o, h_o, w_o = start[i] - - f, h, w = g[1][i] - t_f, t_h, t_w = g[2][i] - seq_f, seq_h, seq_w = f - f_o, h - h_o, w - w_o - seq_len = int(seq_f * seq_h * seq_w) - if seq_len > 0: - if t_f > 0: - factor_f, factor_h, factor_w = (t_f / seq_f).item(), ( - t_h / seq_h).item(), (t_w / seq_w).item() - - if f_o >= 0: - f_sam = np.linspace(f_o.item(), (t_f + f_o).item() - 1, - seq_f).astype(int).tolist() - else: - f_sam = np.linspace(-f_o.item(), - (-t_f - f_o).item() + 1, - seq_f).astype(int).tolist() - h_sam = np.linspace(h_o.item(), (t_h + h_o).item() - 1, - seq_h).astype(int).tolist() - w_sam = np.linspace(w_o.item(), (t_w + w_o).item() - 1, - seq_w).astype(int).tolist() - - assert f_o * f >= 0 and h_o * h >= 0 and w_o * w >= 0 - freqs_0 = freqs[0][f_sam] if f_o >= 0 else freqs[0][ - f_sam].conj() - freqs_0 = freqs_0.view(seq_f, 1, 1, -1) - - freqs_i = torch.cat([ - freqs_0.expand(seq_f, seq_h, seq_w, -1), - freqs[1][h_sam].view(1, seq_h, 1, -1).expand( - seq_f, seq_h, seq_w, -1), - freqs[2][w_sam].view(1, 1, seq_w, -1).expand( - seq_f, seq_h, seq_w, -1), - ], - dim=-1).reshape(seq_len, 1, -1) - elif t_f < 0: - freqs_i = trainable_freqs.unsqueeze(1) - # apply rotary embedding - # precompute multipliers - x_i = torch.view_as_complex( - x[i, seq_bucket[-1]:seq_bucket[-1] + seq_len].to( - torch.float64).reshape(seq_len, n, -1, 2)) - x_i = torch.view_as_real(x_i * freqs_i).flatten(2) - output[i, seq_bucket[-1]:seq_bucket[-1] + seq_len] = x_i - seq_bucket.append(seq_bucket[-1] + seq_len) - return output.float() - - -class RMSNorm(nn.Module): - - def __init__(self, dim, eps=1e-5): - super().__init__() - self.dim = dim - self.eps = eps - self.weight = nn.Parameter(torch.ones(dim)) - - def forward(self, x): - return self._norm(x.float()).type_as(x) * self.weight - - def _norm(self, x): - return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) - - -class LayerNorm(nn.LayerNorm): - - def __init__(self, dim, eps=1e-6, elementwise_affine=False): - super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps) - - def forward(self, x): - return super().forward(x.float()).type_as(x) - - -class SelfAttention(nn.Module): - - def __init__(self, - dim, - num_heads, - window_size=(-1, -1), - qk_norm=True, - eps=1e-6): - assert dim % num_heads == 0 - super().__init__() - self.dim = dim - self.num_heads = num_heads - self.head_dim = dim // num_heads - self.window_size = window_size - self.qk_norm = qk_norm - self.eps = eps - - # layers - self.q = nn.Linear(dim, dim) - self.k = nn.Linear(dim, dim) - self.v = nn.Linear(dim, dim) - self.o = nn.Linear(dim, dim) - self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() - self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity() - - def forward(self, x, seq_lens, grid_sizes, freqs): - b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim - - # query, key, value function - def qkv_fn(x): - q = self.norm_q(self.q(x)).view(b, s, n, d) - k = self.norm_k(self.k(x)).view(b, s, n, d) - v = self.v(x).view(b, s, n, d) - return q, k, v - - q, k, v = qkv_fn(x) - - x = attention( - q=rope_apply(q, grid_sizes, freqs), - k=rope_apply(k, grid_sizes, freqs), - v=v, - k_lens=seq_lens) - - # output - x = x.flatten(2) - x = self.o(x) - return x - - -class SwinSelfAttention(SelfAttention): - - def forward(self, x, seq_lens, grid_sizes, freqs): - b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim - assert b == 1, 'Only support batch_size 1' - - # query, key, value function - def qkv_fn(x): - q = self.norm_q(self.q(x)).view(b, s, n, d) - k = self.norm_k(self.k(x)).view(b, s, n, d) - v = self.v(x).view(b, s, n, d) - return q, k, v - - q, k, v = qkv_fn(x) - - q = rope_apply(q, grid_sizes, freqs) - k = rope_apply(k, grid_sizes, freqs) - T, H, W = grid_sizes[0].tolist() - - q = rearrange(q, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) - k = rearrange(k, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) - v = rearrange(v, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) - - ref_q = q[-1:] - q = q[:-1] - - ref_k = repeat( - k[-1:], "1 s n d -> t s n d", t=k.shape[0] - 1) # t hw n d - k = k[:-1] - k = torch.cat([k[:1], k, k[-1:]]) - k = torch.cat([k[1:-1], k[2:], k[:-2], ref_k], dim=1) # (bt) (3hw) n d - - ref_v = repeat(v[-1:], "1 s n d -> t s n d", t=v.shape[0] - 1) - v = v[:-1] - v = torch.cat([v[:1], v, v[-1:]]) - v = torch.cat([v[1:-1], v[2:], v[:-2], ref_v], dim=1) - - # q: b (t h w) n d - # k: b (t h w) n d - out = attention(q, k, v) - out = torch.cat([out, ref_v[:1]], axis=0) - out = rearrange(out, '(b t) (h w) n d -> b (t h w) n d', t=T, h=H, w=W) - x = out - - # output - x = x.flatten(2) - x = self.o(x) - return x - - -#Fix the reference frame RoPE to 1,H,W. -#Set the current frame RoPE to 1. -#Set the previous frame RoPE to 0. -class CasualSelfAttention(SelfAttention): - - def forward(self, x, seq_lens, grid_sizes, freqs): - shifting = 3 - b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim - assert b == 1, 'Only support batch_size 1' - - # query, key, value function - def qkv_fn(x): - q = self.norm_q(self.q(x)).view(b, s, n, d) - k = self.norm_k(self.k(x)).view(b, s, n, d) - v = self.v(x).view(b, s, n, d) - return q, k, v - - q, k, v = qkv_fn(x) - - T, H, W = grid_sizes[0].tolist() - - q = rearrange(q, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) - k = rearrange(k, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) - v = rearrange(v, 'b (t h w) n d -> (b t) (h w) n d', t=T, h=H, w=W) - - ref_q = q[-1:] - q = q[:-1] - - grid_sizes = torch.tensor([[1, H, W]] * q.shape[0], dtype=torch.long) - start = [[shifting, 0, 0]] * q.shape[0] - q = rope_apply(q, grid_sizes, freqs, start=start) - - ref_k = k[-1:] - grid_sizes = torch.tensor([[1, H, W]], dtype=torch.long) - # start = [[shifting, H, W]] - - start = [[shifting + 10, 0, 0]] - ref_k = rope_apply(ref_k, grid_sizes, freqs, start) - ref_k = repeat( - ref_k, "1 s n d -> t s n d", t=k.shape[0] - 1) # t hw n d - - k = k[:-1] - k = torch.cat([*([k[:1]] * shifting), k]) - cat_k = [] - for i in range(shifting): - cat_k.append(k[i:i - shifting]) - cat_k.append(k[shifting:]) - k = torch.cat(cat_k, dim=1) # (bt) (3hw) n d - - grid_sizes = torch.tensor( - [[shifting + 1, H, W]] * q.shape[0], dtype=torch.long) - k = rope_apply(k, grid_sizes, freqs) - k = torch.cat([k, ref_k], dim=1) - - ref_v = repeat(v[-1:], "1 s n d -> t s n d", t=q.shape[0]) # t hw n d - v = v[:-1] - v = torch.cat([*([v[:1]] * shifting), v]) - cat_v = [] - for i in range(shifting): - cat_v.append(v[i:i - shifting]) - cat_v.append(v[shifting:]) - v = torch.cat(cat_v, dim=1) # (bt) (3hw) n d - v = torch.cat([v, ref_v], dim=1) - - # q: b (t h w) n d - # k: b (t h w) n d - outs = [] - for i in range(q.shape[0]): - out = attention( - q=q[i:i + 1], - k=k[i:i + 1], - v=v[i:i + 1], - window_size=self.window_size) - outs.append(out) - out = torch.cat(outs, dim=0) - out = torch.cat([out, ref_v[:1]], axis=0) - out = rearrange(out, '(b t) (h w) n d -> b (t h w) n d', t=T, h=H, w=W) - x = out - - # output - x = x.flatten(2) - x = self.o(x) - return x - - -class MotionerAttentionBlock(nn.Module): - - def __init__(self, - dim, - ffn_dim, - num_heads, - window_size=(-1, -1), - qk_norm=True, - cross_attn_norm=False, - eps=1e-6, - self_attn_block="SelfAttention"): - super().__init__() - self.dim = dim - self.ffn_dim = ffn_dim - self.num_heads = num_heads - self.window_size = window_size - self.qk_norm = qk_norm - self.cross_attn_norm = cross_attn_norm - self.eps = eps - - # layers - self.norm1 = LayerNorm(dim, eps) - if self_attn_block == "SelfAttention": - self.self_attn = SelfAttention(dim, num_heads, window_size, qk_norm, - eps) - elif self_attn_block == "SwinSelfAttention": - self.self_attn = SwinSelfAttention(dim, num_heads, window_size, - qk_norm, eps) - elif self_attn_block == "CasualSelfAttention": - self.self_attn = CasualSelfAttention(dim, num_heads, window_size, - qk_norm, eps) - - self.norm2 = LayerNorm(dim, eps) - self.ffn = nn.Sequential( - nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'), - nn.Linear(ffn_dim, dim)) - - def forward( - self, - x, - seq_lens, - grid_sizes, - freqs, - ): - # self-attention - y = self.self_attn(self.norm1(x).float(), seq_lens, grid_sizes, freqs) - x = x + y - y = self.ffn(self.norm2(x).float()) - x = x + y - return x - - -class Head(nn.Module): - - def __init__(self, dim, out_dim, patch_size, eps=1e-6): - super().__init__() - self.dim = dim - self.out_dim = out_dim - self.patch_size = patch_size - self.eps = eps - - # layers - out_dim = math.prod(patch_size) * out_dim - self.norm = LayerNorm(dim, eps) - self.head = nn.Linear(dim, out_dim) - - def forward(self, x): - x = self.head(self.norm(x)) - return x - - -class MotionerTransformers(nn.Module, PeftAdapterMixin): - - def __init__( - self, - patch_size=(1, 2, 2), - in_dim=16, - dim=2048, - ffn_dim=8192, - freq_dim=256, - out_dim=16, - num_heads=16, - num_layers=32, - window_size=(-1, -1), - qk_norm=True, - cross_attn_norm=False, - eps=1e-6, - self_attn_block="SelfAttention", - motion_token_num=1024, - enable_tsm=False, - motion_stride=4, - expand_ratio=2, - trainable_token_pos_emb=False, - ): - super().__init__() - self.patch_size = patch_size - self.in_dim = in_dim - self.dim = dim - self.ffn_dim = ffn_dim - self.freq_dim = freq_dim - self.out_dim = out_dim - self.num_heads = num_heads - self.num_layers = num_layers - self.window_size = window_size - self.qk_norm = qk_norm - self.cross_attn_norm = cross_attn_norm - self.eps = eps - - self.enable_tsm = enable_tsm - self.motion_stride = motion_stride - self.expand_ratio = expand_ratio - self.sample_c = self.patch_size[0] - - # embeddings - self.patch_embedding = nn.Conv3d( - in_dim, dim, kernel_size=patch_size, stride=patch_size) - - # blocks - self.blocks = nn.ModuleList([ - MotionerAttentionBlock( - dim, - ffn_dim, - num_heads, - window_size, - qk_norm, - cross_attn_norm, - eps, - self_attn_block=self_attn_block) for _ in range(num_layers) - ]) - - # buffers (don't use register_buffer otherwise dtype will be changed in to()) - assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0 - d = dim // num_heads - self.freqs = torch.cat([ - rope_params(1024, d - 4 * (d // 6)), - rope_params(1024, 2 * (d // 6)), - rope_params(1024, 2 * (d // 6)) - ], - dim=1) - - self.gradient_checkpointing = False - - self.motion_side_len = int(math.sqrt(motion_token_num)) - assert self.motion_side_len**2 == motion_token_num - self.token = nn.Parameter( - torch.zeros(1, motion_token_num, dim).contiguous()) - - self.trainable_token_pos_emb = trainable_token_pos_emb - if trainable_token_pos_emb: - x = torch.zeros([1, motion_token_num, num_heads, d]) - x[..., ::2] = 1 - - gride_sizes = [[ - torch.tensor([0, 0, 0]).unsqueeze(0).repeat(1, 1), - torch.tensor([1, self.motion_side_len, self.motion_side_len]).unsqueeze(0).repeat(1, 1), - torch.tensor([1, self.motion_side_len, self.motion_side_len]).unsqueeze(0).repeat(1, 1), - ]] - token_freqs = rope_apply(x, gride_sizes, self.freqs) - token_freqs = token_freqs[0, :, 0].reshape(motion_token_num, -1, 2) - token_freqs = token_freqs * 0.01 - self.token_freqs = torch.nn.Parameter(token_freqs) - - def after_patch_embedding(self, x): - return x - - def forward( - self, - x, - ): - """ - x: A list of videos each with shape [C, T, H, W]. - t: [B]. - context: A list of text embeddings each with shape [L, C]. - """ - # params - device = self.patch_embedding.weight.device - freqs = self.freqs - if freqs.device != device: - freqs = freqs.to(device) - - if self.trainable_token_pos_emb: - with amp.autocast(dtype=torch.float64): - token_freqs = self.token_freqs.to(torch.float64) - token_freqs = token_freqs / token_freqs.norm( - dim=-1, keepdim=True) - freqs = [freqs, torch.view_as_complex(token_freqs)] - - if self.enable_tsm: - sample_idx = [ - sample_indices( - u.shape[1], - stride=self.motion_stride, - expand_ratio=self.expand_ratio, - c=self.sample_c) for u in x - ] - x = [ - torch.flip(torch.flip(u, [1])[:, idx], [1]) - for idx, u in zip(sample_idx, x) - ] - - # embeddings - x = [self.patch_embedding(u.unsqueeze(0)) for u in x] - x = self.after_patch_embedding(x) - - seq_f, seq_h, seq_w = x[0].shape[-3:] - batch_size = len(x) - if not self.enable_tsm: - grid_sizes = torch.stack([torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) - grid_sizes = [[torch.zeros_like(grid_sizes), grid_sizes, grid_sizes]] - seq_f = 0 - else: - grid_sizes = [] - for idx in sample_idx[0][::-1][::self.sample_c]: - tsm_frame_grid_sizes = [[ - torch.tensor([idx, 0, 0]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([idx + 1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1), - ]] - grid_sizes += tsm_frame_grid_sizes - seq_f = sample_idx[0][-1] + 1 - - x = [u.flatten(2).transpose(1, 2) for u in x] - seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long) - x = torch.cat([u for u in x]) - - batch_size = len(x) - - token_grid_sizes = [[ - torch.tensor([seq_f, 0, 0]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([seq_f + 1, self.motion_side_len,self.motion_side_len]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([1 if not self.trainable_token_pos_emb else -1, seq_h, seq_w]).unsqueeze(0).repeat(batch_size, 1), - ] # 第三行代表rope emb的想要覆盖到的范围 - ] - - grid_sizes = grid_sizes + token_grid_sizes - token_unpatch_grid_sizes = torch.stack([ - torch.tensor([1, 32, 32], dtype=torch.long) - for b in range(batch_size) - ]) - token_len = self.token.shape[1] - token = self.token.clone().repeat(x.shape[0], 1, 1).contiguous() - seq_lens = seq_lens + torch.tensor([t.size(0) for t in token], - dtype=torch.long) - x = torch.cat([x, token], dim=1) - # arguments - kwargs = dict( - seq_lens=seq_lens, - grid_sizes=grid_sizes, - freqs=freqs, - ) - - for idx, block in enumerate(self.blocks): - x = block(x, **kwargs) - # head - out = x[:, -token_len:] - return out - - def unpatchify(self, x, grid_sizes): - c = self.out_dim - out = [] - for u, v in zip(x, grid_sizes.tolist()): - u = u[:math.prod(v)].view(*v, *self.patch_size, c) - u = torch.einsum('fhwpqrc->cfphqwr', u) - u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)]) - out.append(u) - return out - -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 - *args, - **kwargs): - super().__init__(*args, **kwargs) - 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 = torch.tensor( - zip_frame_buckets, dtype=torch.long) - - self.inner_dim = inner_dim - self.num_heads = num_heads - - assert (inner_dim % - num_heads) == 0 and (inner_dim // num_heads) % 2 == 0 - d = inner_dim // num_heads - self.freqs = torch.cat([ - rope_params(1024, d - 4 * (d // 6)), - rope_params(1024, 2 * (d // 6)), - rope_params(1024, 2 * (d // 6)) - ], - dim=1) - self.drop_mode = drop_mode - - def forward(self, motion_latents, add_last_motion=2): - mot = [] - mot_remb = [] - for m in motion_latents: - lat_height, lat_width = m.shape[2], m.shape[3] - padd_lat = torch.zeros(16, self.zip_frame_buckets.sum(), lat_height, - lat_width).to( - device=m.device, dtype=m.dtype) - overlap_frame = min(padd_lat.shape[1], m.shape[1]) - if overlap_frame > 0: - padd_lat[:, -overlap_frame:] = m[:, -overlap_frame:] - - if add_last_motion < 2 and self.drop_mode != "drop": - zero_end_frame = self.zip_frame_buckets[:self.zip_frame_buckets. - __len__() - - add_last_motion - - 1].sum() - padd_lat[:, -zero_end_frame:] = 0 - - padd_lat = padd_lat.unsqueeze(0) - clean_latents_4x, clean_latents_2x, clean_latents_post = padd_lat[:, :, -self.zip_frame_buckets.sum( - ):, :, :].split( - list(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).flatten(2).transpose(1, 2) - clean_latents_4x = self.proj_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 - start_time_id = -(self.zip_frame_buckets[:1].sum()) - end_time_id = start_time_id + self.zip_frame_buckets[0] - grid_sizes = [] if add_last_motion < 2 and self.drop_mode == "drop" else \ - [ - [torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), - torch.tensor([end_time_id, lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), - torch.tensor([self.zip_frame_buckets[0], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), ] - ] - - start_time_id = -(self.zip_frame_buckets[:2].sum()) - end_time_id = start_time_id + self.zip_frame_buckets[1] // 2 - grid_sizes_2x = [] if add_last_motion < 1 and self.drop_mode == "drop" else \ - [[ - torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), - torch.tensor([end_time_id, lat_height // 4, lat_width // 4]).unsqueeze(0).repeat(1, 1), - torch.tensor([self.zip_frame_buckets[1], lat_height // 2, lat_width // 2]).unsqueeze(0).repeat(1, 1), - ]] - - start_time_id = -(self.zip_frame_buckets[:3].sum()) - end_time_id = start_time_id + self.zip_frame_buckets[2] // 4 - grid_sizes_4x = [[ - torch.tensor([start_time_id, 0, 0]).unsqueeze(0).repeat(1, 1), - torch.tensor([end_time_id, lat_height // 8, lat_width // 8]).unsqueeze(0).repeat(1, 1), - torch.tensor([self.zip_frame_buckets[2], lat_height // 2, lat_width // 2 - ]).unsqueeze(0).repeat(1, 1), - ]] - - grid_sizes = grid_sizes + grid_sizes_2x + grid_sizes_4x - - motion_rope_emb = rope_precompute( - motion_lat.detach().view(1, motion_lat.shape[1], self.num_heads, - self.inner_dim // self.num_heads), - grid_sizes, - self.freqs, - start=None) - - mot.append(motion_lat) - mot_remb.append(motion_rope_emb) - return mot, mot_remb - - -def sample_indices(N, stride, expand_ratio, c): - indices = [] - current_start = 0 - - while current_start < N: - bucket_width = int(stride * (expand_ratio**(len(indices) / stride))) - - interval = int(bucket_width / stride * c) - current_end = min(N, current_start + bucket_width) - bucket_samples = [] - for i in range(current_end - 1, current_start - 1, -interval): - for near in range(c): - bucket_samples.append(i - near) - - indices += bucket_samples[::-1] - current_start += bucket_width - - return indices - diff --git a/wanvideo/modules/s2v/s2v_utils.py b/wanvideo/modules/s2v/s2v_utils.py deleted file mode 100644 index 68644a2..0000000 --- a/wanvideo/modules/s2v/s2v_utils.py +++ /dev/null @@ -1,70 +0,0 @@ -# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. -import numpy as np -import torch - - -def rope_precompute(x, grid_sizes, freqs, start=None): - b, s, n, c = x.size(0), x.size(1), x.size(2), x.size(3) // 2 - - # split freqs - if type(freqs) is list: - trainable_freqs = freqs[1] - freqs = freqs[0] - freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1) - - # loop over samples - output = torch.view_as_complex(x.detach().reshape(b, s, n, -1, - 2).to(torch.float64)) - seq_bucket = [0] - if not type(grid_sizes) is list: - grid_sizes = [grid_sizes] - for g in grid_sizes: - if not type(g) is list: - g = [torch.zeros_like(g), g] - batch_size = g[0].shape[0] - for i in range(batch_size): - if start is None: - f_o, h_o, w_o = g[0][i] - else: - f_o, h_o, w_o = start[i] - - f, h, w = g[1][i] - t_f, t_h, t_w = g[2][i] - seq_f, seq_h, seq_w = f - f_o, h - h_o, w - w_o - seq_len = int(seq_f * seq_h * seq_w) - if seq_len > 0: - if t_f > 0: - factor_f, factor_h, factor_w = (t_f / seq_f).item(), ( - t_h / seq_h).item(), (t_w / seq_w).item() - # Generate a list of seq_f integers starting from f_o and ending at math.ceil(factor_f * seq_f.item() + f_o.item()) - if f_o >= 0: - f_sam = np.linspace(f_o.item(), (t_f + f_o).item() - 1, - seq_f).astype(int).tolist() - else: - f_sam = np.linspace(-f_o.item(), - (-t_f - f_o).item() + 1, - seq_f).astype(int).tolist() - h_sam = np.linspace(h_o.item(), (t_h + h_o).item() - 1, - seq_h).astype(int).tolist() - w_sam = np.linspace(w_o.item(), (t_w + w_o).item() - 1, - seq_w).astype(int).tolist() - - assert f_o * f >= 0 and h_o * h >= 0 and w_o * w >= 0 - freqs_0 = freqs[0][f_sam] if f_o >= 0 else freqs[0][ - f_sam].conj() - freqs_0 = freqs_0.view(seq_f, 1, 1, -1) - - freqs_i = torch.cat([ - freqs_0.expand(seq_f, seq_h, seq_w, -1), - freqs[1][h_sam].view(1, seq_h, 1, -1).expand( - seq_f, seq_h, seq_w, -1), - freqs[2][w_sam].view(1, 1, seq_w, -1).expand( - seq_f, seq_h, seq_w, -1), - ], - dim=-1).reshape(seq_len, 1, -1) - elif t_f < 0: - freqs_i = trainable_freqs.unsqueeze(1) - # apply rotary embedding - output[i, seq_bucket[-1]:seq_bucket[-1] + seq_len] = freqs_i - seq_bucket.append(seq_bucket[-1] + seq_len) - return output From a37b12235fbc6f70dcabdba7baa6a750d9f1542d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 20:51:08 +0300 Subject: [PATCH 12/31] Add loudness norm node --- nodes_utility.py | 45 +++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 43 insertions(+), 2 deletions(-) diff --git a/nodes_utility.py b/nodes_utility.py index 9af2016..bb4e408 100644 --- a/nodes_utility.py +++ b/nodes_utility.py @@ -409,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, @@ -417,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", @@ -427,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" } \ No newline at end of file From 5266959a93021310cd0698a6d06680206027eb36 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 21:41:34 +0300 Subject: [PATCH 13/31] Create wanvideo2_2_S2V_context_window_testing.json --- ...anvideo2_2_S2V_context_window_testing.json | 2167 +++++++++++++++++ 1 file changed, 2167 insertions(+) create mode 100644 s2v/wanvideo2_2_S2V_context_window_testing.json diff --git a/s2v/wanvideo2_2_S2V_context_window_testing.json b/s2v/wanvideo2_2_S2V_context_window_testing.json new file mode 100644 index 0000000..16e0a8e --- /dev/null +++ b/s2v/wanvideo2_2_S2V_context_window_testing.json @@ -0,0 +1,2167 @@ +{ + "id": "c6e410bc-5e2c-460b-ae81-c91b6094fbb1", + "revision": 0, + "last_node_id": 100, + "last_link_id": 155, + "nodes": [ + { + "id": 44, + "type": "Note", + "pos": [ + -710, + -710 + ], + "size": [ + 303.0501403808594, + 88 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "If you have Triton installed, connect this for ~30% speed increase" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 38, + "type": "WanVideoVAELoader", + "pos": [ + 1988.66015625, + -572.3654174804688 + ], + "size": [ + 315, + 82 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [ + { + "name": "compile_args", + "shape": 7, + "type": "WANCOMPILEARGS", + "link": null + } + ], + "outputs": [ + { + "name": "vae", + "type": "WANVAE", + "slot_index": 0, + "links": [ + 43, + 81 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoVAELoader" + }, + "widgets_values": [ + "wanvideo\\Wan2_1_VAE_bf16.safetensors", + "bf16" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 56, + "type": "WanVideoSetBlockSwap", + "pos": [ + 882.7855834960938, + -362.668701171875 + ], + "size": [ + 201.76815795898438, + 46 + ], + "flags": {}, + "order": 25, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 62 + }, + { + "name": "block_swap_args", + "shape": 7, + "type": "BLOCKSWAPARGS", + "link": 58 + } + ], + "outputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "links": [ + 60 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoSetBlockSwap" + }, + "widgets_values": [], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 42, + "type": "Note", + "pos": [ + -340.0147399902344, + -394.8644104003906 + ], + "size": [ + 312.98052978515625, + 92.32489013671875 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "Adjust the blocks to swap based on your VRAM, this is a tradeoff between speed and memory usage." + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 36, + "type": "Note", + "pos": [ + 110, + -630 + ], + "size": [ + 374.3061828613281, + 171.9547576904297 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "fp_16_fast enables \"Full FP16 Accmumulation in FP16 GEMMs\" feature available in the very latest pytorch nightly, this is around 20% speed boost. \n\nSageattn if you have it installed can be used for almost double inference speed at higher resolutions\n\nRadial attention is even faster but has worst quality, it should be used along with Set Radial Attention node to control which steps/blocks it's applied on to balance quality and speed." + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 62, + "type": "PreviewAny", + "pos": [ + 546.9254150390625, + -252.02532958984375 + ], + "size": [ + 210, + 88 + ], + "flags": {}, + "order": 23, + "mode": 0, + "inputs": [ + { + "name": "source", + "type": "*", + "link": 65 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "PreviewAny" + }, + "widgets_values": [] + }, + { + "id": 28, + "type": "WanVideoDecode", + "pos": [ + 1994.2247314453125, + -394.9518737792969 + ], + "size": [ + 315, + 198 + ], + "flags": {}, + "order": 31, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "WANVAE", + "link": 43 + }, + { + "name": "samples", + "type": "LATENT", + "link": 105 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 77 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoDecode" + }, + "widgets_values": [ + false, + 272, + 272, + 144, + 128, + "default" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 73, + "type": "LoadImage", + "pos": [ + 609.2274169921875, + 131.5489044189453 + ], + "size": [ + 274.080078125, + 314 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 83 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "2b.jpg", + "image" + ] + }, + { + "id": 67, + "type": "WanVideoTextEncodeCached", + "pos": [ + 46.950748443603516, + -20.68189239501953 + ], + "size": [ + 459.45745849609375, + 393.8887939453125 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "extender_args", + "shape": 7, + "type": "WANVIDEOPROMPTEXTENDER_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "links": [ + 71 + ] + }, + { + "name": "negative_text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "links": null + }, + { + "name": "positive_prompt", + "type": "STRING", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "a1ca0985ec120ff97e34676a64de19c99767bbd4", + "Node name for S&R": "WanVideoTextEncodeCached" + }, + "widgets_values": [ + "umt5-xxl-enc-bf16.safetensors", + "bf16", + "a woman is singing passionately", + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", + "disabled", + true, + "gpu" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 61, + "type": "MarkdownNote", + "pos": [ + -128.9186248779297, + -1117.54248046875 + ], + "size": [ + 510.2661437988281, + 245.62203979492188 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "Models:\n\n[https://huggingface.co/Kijai/WanVideo_comfy_fp8_scaled/blob/main/T2V/Wan2_1-T2V-14B_fp8_e4m3fn_scaled_KJ.safetensors](https://huggingface.co/Kijai/WanVideo_comfy_fp8_scaled/blob/main/T2V/Wan2_1-T2V-14B_fp8_e4m3fn_scaled_KJ.safetensors)\n\nIf you want to use torch compile on GPUs prior to 4000 series:\n\n[https://huggingface.co/Kijai/WanVideo_comfy_fp8_scaled/blob/main/T2V/Wan2_1-T2V-14B_fp8_e5m2_scaled_KJ.safetensors](https://huggingface.co/Kijai/WanVideo_comfy_fp8_scaled/blob/main/T2V/Wan2_1-T2V-14B_fp8_e5m2_scaled_KJ.safetensors)\n\nLoRA:\n\n[https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Lightx2v/lightx2v_T2V_14B_cfg_step_distill_v2_lora_rank64_bf16.safetensors](https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Lightx2v/lightx2v_T2V_14B_cfg_step_distill_v2_lora_rank64_bf16.safetensors)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 72, + "type": "WanVideoEncode", + "pos": [ + 1256.0252685546875, + -81.86726379394531 + ], + "size": [ + 270, + 242 + ], + "flags": {}, + "order": 20, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "WANVAE", + "link": 81 + }, + { + "name": "image", + "type": "IMAGE", + "link": 84 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "links": [ + 92, + 128 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "63d4b6aadaae543f96d101788122563f3e2ba0c8", + "Node name for S&R": "WanVideoEncode" + }, + "widgets_values": [ + false, + 272, + 272, + 144, + 128, + 0, + 1 + ] + }, + { + "id": 22, + "type": "WanVideoModelLoader", + "pos": [ + 10, + -390 + ], + "size": [ + 477.4410095214844, + 314 + ], + "flags": {}, + "order": 18, + "mode": 0, + "inputs": [ + { + "name": "compile_args", + "shape": 7, + "type": "WANCOMPILEARGS", + "link": 129 + }, + { + "name": "block_swap_args", + "shape": 7, + "type": "BLOCKSWAPARGS", + "link": null + }, + { + "name": "lora", + "shape": 7, + "type": "WANVIDLORA", + "link": null + }, + { + "name": "vram_management_args", + "shape": 7, + "type": "VRAM_MANAGEMENTARGS", + "link": null + }, + { + "name": "extra_model", + "shape": 7, + "type": "VACEPATH", + "link": null + }, + { + "name": "fantasytalking_model", + "shape": 7, + "type": "FANTASYTALKINGMODEL", + "link": null + }, + { + "name": "multitalk_model", + "shape": 7, + "type": "MULTITALKMODEL", + "link": null + }, + { + "name": "fantasyportrait_model", + "shape": 7, + "type": "FANTASYPORTRAITMODEL", + "link": null + }, + { + "name": "vace_model", + "shape": 7, + "type": "VACEPATH", + "link": null + } + ], + "outputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "slot_index": 0, + "links": [ + 61, + 65 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoModelLoader" + }, + "widgets_values": [ + "WanVideo\\S2V\\Wan2_2-S2V-14B_fp8_e4m3fn_scaled_KJ.safetensors", + "fp16_fast", + "fp8_e4m3fn_scaled", + "offload_device", + "sageattn" + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 35, + "type": "WanVideoTorchCompileSettings", + "pos": [ + -390, + -710 + ], + "size": [ + 390.5999755859375, + 202 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "torch_compile_args", + "type": "WANCOMPILEARGS", + "slot_index": 0, + "links": [ + 129 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoTorchCompileSettings" + }, + "widgets_values": [ + "inductor", + false, + "default", + false, + 64, + true, + 128 + ] + }, + { + "id": 58, + "type": "WanVideoSetLoRAs", + "pos": [ + 630.5015869140625, + -367.1865234375 + ], + "size": [ + 174.53378295898438, + 46 + ], + "flags": {}, + "order": 22, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 61 + }, + { + "name": "lora", + "shape": 7, + "type": "WANVIDLORA", + "link": 64 + } + ], + "outputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "links": [ + 62 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoSetLoRAs" + }, + "widgets_values": [], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 83, + "type": "WanVideoContextOptions", + "pos": [ + 1290.159423828125, + -496.63787841796875 + ], + "size": [ + 275.783203125, + 202 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "reference_latent", + "shape": 7, + "type": "LATENT", + "link": null + } + ], + "outputs": [ + { + "name": "context_options", + "type": "WANVIDCONTEXT", + "links": [ + 121 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "3c79851230c9ab042f52c8e176f349ecc3e51a64", + "Node name for S&R": "WanVideoContextOptions" + }, + "widgets_values": [ + "uniform_standard", + 73, + 4, + 16, + true, + false, + "linear" + ] + }, + { + "id": 77, + "type": "InsertLatentToIndexed", + "pos": [ + 2088.097412109375, + -69.78427124023438 + ], + "size": [ + 270, + 78 + ], + "flags": {}, + "order": 30, + "mode": 0, + "inputs": [ + { + "name": "source", + "type": "LATENT", + "link": 92 + }, + { + "name": "destination", + "type": "LATENT", + "link": 93 + } + ], + "outputs": [ + { + "name": "LATENT", + "type": "LATENT", + "links": [] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "InsertLatentToIndexed" + }, + "widgets_values": [ + 0 + ] + }, + { + "id": 74, + "type": "ImageResizeKJv2", + "pos": [ + 1009.1867065429688, + 210.0146942138672 + ], + "size": [ + 270, + 336 + ], + "flags": {}, + "order": 17, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 83 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 84 + ] + }, + { + "name": "width", + "type": "INT", + "links": [ + 85 + ] + }, + { + "name": "height", + "type": "INT", + "links": [ + 86 + ] + }, + { + "name": "mask", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "ImageResizeKJv2" + }, + "widgets_values": [ + 960, + 640, + "lanczos", + "crop", + "0, 0, 0", + "center", + 2, + "cpu" + ] + }, + { + "id": 39, + "type": "WanVideoBlockSwap", + "pos": [ + 775.6461791992188, + -256.3999328613281 + ], + "size": [ + 315, + 202 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "block_swap_args", + "type": "BLOCKSWAPARGS", + "slot_index": 0, + "links": [ + 58 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoBlockSwap" + }, + "widgets_values": [ + 25, + false, + false, + true, + 0, + 1, + false + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 70, + "type": "GetImageSizeAndCount", + "pos": [ + 2367.16748046875, + -395.861572265625 + ], + "size": [ + 190.86483764648438, + 86 + ], + "flags": {}, + "order": 32, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 77 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 100 + ] + }, + { + "label": "960 width", + "name": "width", + "type": "INT", + "links": null + }, + { + "label": "640 height", + "name": "height", + "type": "INT", + "links": null + }, + { + "label": "601 count", + "name": "count", + "type": "INT", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "GetImageSizeAndCount" + }, + "widgets_values": [] + }, + { + "id": 60, + "type": "WanVideoLoraSelectMulti", + "pos": [ + 497.81256103515625, + -824.0181274414062 + ], + "size": [ + 614.6506958007812, + 342 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "prev_lora", + "shape": 7, + "type": "WANVIDLORA", + "link": null + }, + { + "name": "blocks", + "shape": 7, + "type": "SELECTEDBLOCKS", + "link": null + } + ], + "outputs": [ + { + "name": "lora", + "type": "WANVIDLORA", + "links": [ + 64 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoLoraSelectMulti" + }, + "widgets_values": [ + "WanVideo\\Lightx2v\\lightx2v_T2V_14B_cfg_step_distill_v2_lora_rank64_bf16_.safetensors", + 1.5, + "none", + 1, + "none", + 1, + "none", + 1, + "none", + 1, + false, + false + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 27, + "type": "WanVideoSampler", + "pos": [ + 1616.490966796875, + -391.5707092285156 + ], + "size": [ + 315, + 900.6666870117188 + ], + "flags": {}, + "order": 29, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 60 + }, + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "link": 91 + }, + { + "name": "text_embeds", + "shape": 7, + "type": "WANVIDEOTEXTEMBEDS", + "link": 71 + }, + { + "name": "samples", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "feta_args", + "shape": 7, + "type": "FETAARGS", + "link": null + }, + { + "name": "context_options", + "shape": 7, + "type": "WANVIDCONTEXT", + "link": 121 + }, + { + "name": "cache_args", + "shape": 7, + "type": "CACHEARGS", + "link": null + }, + { + "name": "flowedit_args", + "shape": 7, + "type": "FLOWEDITARGS", + "link": null + }, + { + "name": "slg_args", + "shape": 7, + "type": "SLGARGS", + "link": null + }, + { + "name": "loop_args", + "shape": 7, + "type": "LOOPARGS", + "link": null + }, + { + "name": "experimental_args", + "shape": 7, + "type": "EXPERIMENTALARGS", + "link": null + }, + { + "name": "sigmas", + "shape": 7, + "type": "SIGMAS", + "link": null + }, + { + "name": "unianimate_poses", + "shape": 7, + "type": "UNIANIMATE_POSE", + "link": null + }, + { + "name": "fantasytalking_embeds", + "shape": 7, + "type": "FANTASYTALKING_EMBEDS", + "link": null + }, + { + "name": "uni3c_embeds", + "shape": 7, + "type": "UNI3C_EMBEDS", + "link": null + }, + { + "name": "multitalk_embeds", + "shape": 7, + "type": "MULTITALK_EMBEDS", + "link": null + }, + { + "name": "freeinit_args", + "shape": 7, + "type": "FREEINITARGS", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "slot_index": 0, + "links": [ + 93, + 105 + ] + }, + { + "name": "denoised_samples", + "type": "LATENT", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoSampler" + }, + "widgets_values": [ + 6, + 1, + 4, + 45, + "fixed", + true, + "lcm", + 0, + 1, + false, + "comfy", + 0, + -1, + false + ] + }, + { + "id": 94, + "type": "VHS_LoadAudio", + "pos": [ + -59.62910079956055, + 496.306396484375 + ], + "size": [ + 415.6340026855469, + 126 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "audio", + "type": "AUDIO", + "links": [ + 148, + 153, + 154 + ] + }, + { + "name": "duration", + "type": "FLOAT", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "8e4d79471bf1952154768e8435a9300077b534fa", + "Node name for S&R": "VHS_LoadAudio" + }, + "widgets_values": { + "audio_file": "input/weightoftheworld2.mp4", + "seek_seconds": 0, + "duration": 0 + } + }, + { + "id": 66, + "type": "LoadAudio", + "pos": [ + -40.06507873535156, + 719.501953125 + ], + "size": [ + 361.2844543457031, + 155.81912231445312 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "AUDIO", + "type": "AUDIO", + "links": [] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "LoadAudio" + }, + "widgets_values": [ + "NieR_ Automata - _Weight of the World_ ENG VER. by Lizz Robinett [CyOSTbel3AM].mp3", + null, + null + ] + }, + { + "id": 81, + "type": "MelBandRoFormerModelLoader", + "pos": [ + 563.452392578125, + 647.8037109375 + ], + "size": [ + 316.2164001464844, + 58 + ], + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "model", + "type": "MELROFORMERMODEL", + "links": [ + 106 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-MelBandRoFormer", + "ver": "b40e263224778ec417114d91d8b3b39934e30de5", + "Node name for S&R": "MelBandRoFormerModelLoader" + }, + "widgets_values": [ + "MelBandRoFormer\\MelBandRoformer_fp16.safetensors" + ] + }, + { + "id": 82, + "type": "MelBandRoFormerSampler", + "pos": [ + 555.0003662109375, + 770.088623046875 + ], + "size": [ + 222.73397827148438, + 46 + ], + "flags": {}, + "order": 19, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MELROFORMERMODEL", + "link": 106 + }, + { + "name": "audio", + "type": "AUDIO", + "link": 154 + } + ], + "outputs": [ + { + "name": "vocals", + "type": "AUDIO", + "links": [ + 149 + ] + }, + { + "name": "instruments", + "type": "AUDIO", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-MelBandRoFormer", + "ver": "b40e263224778ec417114d91d8b3b39934e30de5", + "Node name for S&R": "MelBandRoFormerSampler" + }, + "widgets_values": [] + }, + { + "id": 98, + "type": "NormalizeAudioLoudness", + "pos": [ + 834.3135986328125, + 760.9041137695312 + ], + "size": [ + 270, + 58 + ], + "flags": {}, + "order": 24, + "mode": 0, + "inputs": [ + { + "name": "audio", + "type": "AUDIO", + "link": 149 + } + ], + "outputs": [ + { + "name": "audio", + "type": "AUDIO", + "links": [ + 150 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "90c3bbb6c2e4ff5e05305e765d007d5e58428ce4", + "Node name for S&R": "NormalizeAudioLoudness" + }, + "widgets_values": [ + -23 + ] + }, + { + "id": 97, + "type": "VHS_VideoCombine", + "pos": [ + 3319.513671875, + 193.33419799804688 + ], + "size": [ + 940.8292846679688, + 961.8861694335938 + ], + "flags": {}, + "order": 35, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 155 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": 148 + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "330bce6c3c0d47ebdedcc0348d9ab355707b7523", + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "WanVideo2_2_S2V", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "trim_to_audio": false, + "pingpong": false, + "save_output": false, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "WanVideo2_2_S2V_00012-audio.mp4", + "subfolder": "", + "type": "temp", + "format": "video/h264-mp4", + "frame_rate": 16, + "workflow": "WanVideo2_2_S2V_00012.png", + "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo2_2_S2V_00012-audio.mp4" + } + } + } + }, + { + "id": 30, + "type": "VHS_VideoCombine", + "pos": [ + 3649.713623046875, + -868.4019165039062 + ], + "size": [ + 940.8292846679688, + 961.8861694335938 + ], + "flags": {}, + "order": 36, + "mode": 2, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 146 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": 153 + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "330bce6c3c0d47ebdedcc0348d9ab355707b7523", + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 32, + "loop_count": 0, + "filename_prefix": "WanVideo2_2_S2V", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "trim_to_audio": false, + "pingpong": false, + "save_output": false, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "WanVideo2_2_S2V_00013-audio.mp4", + "subfolder": "", + "type": "temp", + "format": "video/h264-mp4", + "frame_rate": 32, + "workflow": "WanVideo2_2_S2V_00013.png", + "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo2_2_S2V_00013-audio.mp4" + } + } + } + }, + { + "id": 96, + "type": "GIMMVFI_interpolate", + "pos": [ + 3256.453369140625, + -512.8270263671875 + ], + "size": [ + 270, + 174 + ], + "flags": {}, + "order": 34, + "mode": 2, + "inputs": [ + { + "name": "gimmvfi_model", + "type": "GIMMVIF_MODEL", + "link": 144 + }, + { + "name": "images", + "type": "IMAGE", + "link": 145 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 146 + ] + }, + { + "name": "flow_tensors", + "type": "IMAGE", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-GIMM-VFI", + "ver": "4c9a3123762af85e7c796e41737da0b70c75d72d", + "Node name for S&R": "GIMMVFI_interpolate" + }, + "widgets_values": [ + 1, + 2, + 0, + "fixed", + false + ] + }, + { + "id": 95, + "type": "DownloadAndLoadGIMMVFIModel", + "pos": [ + 3262.200927734375, + -713.7930908203125 + ], + "size": [ + 339.4301452636719, + 106 + ], + "flags": {}, + "order": 14, + "mode": 2, + "inputs": [], + "outputs": [ + { + "name": "gimmvfi_model", + "type": "GIMMVIF_MODEL", + "links": [ + 144 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-GIMM-VFI", + "ver": "4c9a3123762af85e7c796e41737da0b70c75d72d", + "Node name for S&R": "DownloadAndLoadGIMMVFIModel" + }, + "widgets_values": [ + "gimmvfi_r_arb_lpips_fp32.safetensors", + "fp16", + false + ] + }, + { + "id": 80, + "type": "VHS_SplitImages", + "pos": [ + 2397.607666015625, + -214.64483642578125 + ], + "size": [ + 210, + 118 + ], + "flags": {}, + "order": 33, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 100 + } + ], + "outputs": [ + { + "name": "IMAGE_A", + "type": "IMAGE", + "links": null + }, + { + "name": "A_count", + "type": "INT", + "links": null + }, + { + "name": "IMAGE_B", + "type": "IMAGE", + "links": [ + 145, + 155 + ] + }, + { + "name": "B_count", + "type": "INT", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "8e4d79471bf1952154768e8435a9300077b534fa", + "Node name for S&R": "VHS_SplitImages" + }, + "widgets_values": { + "split_index": 3 + } + }, + { + "id": 63, + "type": "WanVideoAddAudioEmbeds", + "pos": [ + 2055.220703125, + -847.4767456054688 + ], + "size": [ + 254.5941619873047, + 122 + ], + "flags": {}, + "order": 27, + "mode": 0, + "inputs": [ + { + "name": "embeds", + "type": "WANVIDIMAGE_EMBEDS", + "link": 66 + }, + { + "name": "audio_encoder_output", + "type": "AUDIO_ENCODER_OUTPUT", + "link": 68 + }, + { + "name": "ref_latent", + "shape": 7, + "type": "LATENT", + "link": 128 + }, + { + "name": "frames", + "type": "INT", + "widget": { + "name": "frames" + }, + "link": 96 + } + ], + "outputs": [ + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "links": [ + 75, + 91 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "a1ca0985ec120ff97e34676a64de19c99767bbd4", + "Node name for S&R": "WanVideoAddAudioEmbeds" + }, + "widgets_values": [ + 601, + 1 + ] + }, + { + "id": 69, + "type": "PreviewAny", + "pos": [ + 2441.904052734375, + -1017.1683349609375 + ], + "size": [ + 490.69390869140625, + 355.7907409667969 + ], + "flags": {}, + "order": 28, + "mode": 0, + "inputs": [ + { + "name": "source", + "type": "*", + "link": 75 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "PreviewAny" + }, + "widgets_values": [] + }, + { + "id": 37, + "type": "WanVideoEmptyEmbeds", + "pos": [ + 1686.0142822265625, + -1071.16162109375 + ], + "size": [ + 315, + 126 + ], + "flags": {}, + "order": 21, + "mode": 0, + "inputs": [ + { + "name": "control_embeds", + "shape": 7, + "type": "WANVIDIMAGE_EMBEDS", + "link": null + }, + { + "name": "extra_latents", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "width", + "type": "INT", + "widget": { + "name": "width" + }, + "link": 85 + }, + { + "name": "height", + "type": "INT", + "widget": { + "name": "height" + }, + "link": 86 + }, + { + "name": "num_frames", + "type": "INT", + "widget": { + "name": "num_frames" + }, + "link": 79 + } + ], + "outputs": [ + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "links": [ + 66 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoEmptyEmbeds" + }, + "widgets_values": [ + 832, + 480, + 601 + ] + }, + { + "id": 71, + "type": "PrimitiveNode", + "pos": [ + 1682.593505859375, + -863.1051635742188 + ], + "size": [ + 210, + 82 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "widget": { + "name": "num_frames" + }, + "links": [ + 79, + 96 + ] + } + ], + "title": "num_frames", + "properties": { + "Run widget replace on values": false + }, + "widgets_values": [ + 601, + "fixed" + ] + }, + { + "id": 64, + "type": "AudioEncoderEncode", + "pos": [ + 1198.93212890625, + 778.826904296875 + ], + "size": [ + 285.087890625, + 46 + ], + "flags": {}, + "order": 26, + "mode": 0, + "inputs": [ + { + "name": "audio_encoder", + "type": "AUDIO_ENCODER", + "link": 69 + }, + { + "name": "audio", + "type": "AUDIO", + "link": 150 + } + ], + "outputs": [ + { + "name": "AUDIO_ENCODER_OUTPUT", + "type": "AUDIO_ENCODER_OUTPUT", + "links": [ + 68 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "AudioEncoderEncode" + }, + "widgets_values": [] + }, + { + "id": 65, + "type": "AudioEncoderLoader", + "pos": [ + 1136.3577880859375, + 641.1036376953125 + ], + "size": [ + 346.756103515625, + 58 + ], + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "AUDIO_ENCODER", + "type": "AUDIO_ENCODER", + "links": [ + 69 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "AudioEncoderLoader" + }, + "widgets_values": [ + "wav2vec_xlsr_53_english_fp32.safetensors" + ] + } + ], + "links": [ + [ + 43, + 38, + 0, + 28, + 0, + "VAE" + ], + [ + 58, + 39, + 0, + 56, + 1, + "BLOCKSWAPARGS" + ], + [ + 60, + 56, + 0, + 27, + 0, + "WANVIDEOMODEL" + ], + [ + 61, + 22, + 0, + 58, + 0, + "WANVIDEOMODEL" + ], + [ + 62, + 58, + 0, + 56, + 0, + "WANVIDEOMODEL" + ], + [ + 64, + 60, + 0, + 58, + 1, + "WANVIDLORA" + ], + [ + 65, + 22, + 0, + 62, + 0, + "*" + ], + [ + 66, + 37, + 0, + 63, + 0, + "WANVIDIMAGE_EMBEDS" + ], + [ + 68, + 64, + 0, + 63, + 1, + "AUDIO_ENCODER_OUTPUT" + ], + [ + 69, + 65, + 0, + 64, + 0, + "AUDIO_ENCODER" + ], + [ + 71, + 67, + 0, + 27, + 2, + "WANVIDEOTEXTEMBEDS" + ], + [ + 75, + 63, + 0, + 69, + 0, + "*" + ], + [ + 77, + 28, + 0, + 70, + 0, + "IMAGE" + ], + [ + 79, + 71, + 0, + 37, + 4, + "INT" + ], + [ + 81, + 38, + 0, + 72, + 0, + "WANVAE" + ], + [ + 83, + 73, + 0, + 74, + 0, + "IMAGE" + ], + [ + 84, + 74, + 0, + 72, + 1, + "IMAGE" + ], + [ + 85, + 74, + 1, + 37, + 2, + "INT" + ], + [ + 86, + 74, + 2, + 37, + 3, + "INT" + ], + [ + 91, + 63, + 0, + 27, + 1, + "WANVIDIMAGE_EMBEDS" + ], + [ + 92, + 72, + 0, + 77, + 0, + "LATENT" + ], + [ + 93, + 27, + 0, + 77, + 1, + "LATENT" + ], + [ + 96, + 71, + 0, + 63, + 3, + "INT" + ], + [ + 100, + 70, + 0, + 80, + 0, + "IMAGE" + ], + [ + 105, + 27, + 0, + 28, + 1, + "LATENT" + ], + [ + 106, + 81, + 0, + 82, + 0, + "MELROFORMERMODEL" + ], + [ + 121, + 83, + 0, + 27, + 5, + "WANVIDCONTEXT" + ], + [ + 128, + 72, + 0, + 63, + 2, + "LATENT" + ], + [ + 129, + 35, + 0, + 22, + 0, + "WANCOMPILEARGS" + ], + [ + 144, + 95, + 0, + 96, + 0, + "GIMMVIF_MODEL" + ], + [ + 145, + 80, + 2, + 96, + 1, + "IMAGE" + ], + [ + 146, + 96, + 0, + 30, + 0, + "IMAGE" + ], + [ + 148, + 94, + 0, + 97, + 1, + "AUDIO" + ], + [ + 149, + 82, + 0, + 98, + 0, + "AUDIO" + ], + [ + 150, + 98, + 0, + 64, + 1, + "AUDIO" + ], + [ + 153, + 94, + 0, + 30, + 1, + "AUDIO" + ], + [ + 154, + 94, + 0, + 82, + 1, + "AUDIO" + ], + [ + 155, + 80, + 2, + 97, + 0, + "IMAGE" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 0.505447028499326, + "offset": [ + 740.042928663884, + 1057.007052457307 + ] + }, + "frontendVersion": "1.26.6", + "node_versions": { + "ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd", + "comfy-core": "0.3.26", + "ComfyUI-VideoHelperSuite": "0a75c7958fe320efcb052f1d9f8451fd20c730a8" + }, + "VHS_latentpreview": true, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true + }, + "version": 0.4 +} \ No newline at end of file From a5621b87391013155b4f688fbe01dba10a8104aa Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 27 Aug 2025 22:43:28 +0300 Subject: [PATCH 14/31] Add pose input --- nodes.py | 96 +++++++++++++++++++++++++++++++++++---- nodes_model_loading.py | 2 +- s2v/nodes.py | 53 +++++++++++---------- wanvideo/modules/model.py | 13 ++++-- 4 files changed, 126 insertions(+), 38 deletions(-) diff --git a/nodes.py b/nodes.py index dd2810b..50a570e 100644 --- a/nodes.py +++ b/nodes.py @@ -2228,17 +2228,23 @@ class WanVideoSampler: 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_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"]].to(device, dtype) s2v_audio_scale = s2v_audio_embeds["audio_scale"] - s2v_ref_latent = s2v_audio_embeds["ref_latent"].to(device, dtype) if "ref_latent" in s2v_audio_embeds else None - s2v_ref_motion = s2v_audio_embeds["ref_motion"].to(device, dtype) if "ref_motion" in s2v_audio_embeds else None - s2v_audio_input = s2v_audio_input[..., 0:image_embeds["num_frames"]] + 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_num_repeat = s2v_audio_embeds.get("num_repeat", 1) vae = image_embeds.get("vae", None) framepack = False - #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) # vid2vid noise_mask=original_image=None @@ -2700,7 +2706,8 @@ class WanVideoSampler: "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_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 } batch_size = 1 @@ -3664,7 +3671,78 @@ class WanVideoSampler: except: pass return {"video": gen_video_samples.permute(1, 2, 3, 0)}, - + elif framepack: + framepack_out = [] + ref_motion_image = None + motion_frames = 5 + infer_frames = image_embeds["num_frames"] + + for r in range(s2v_num_repeat): + if ref_motion_image is not None: + if ref_motion_image.shape[0] > 73: + ref_motion_image = ref_motion_image[-73:] + + if ref_motion_image.shape[0] < 73: + ref = torch.ones([73, ref_motion_image.shape[1], ref_motion_image.shape[2], 3]) * 0.5 + ref[-ref_motion_image.shape[0]:] = ref_motion_image + ref_motion_image = ref + + vae.to(device) + ref_motion = vae.encode(ref_motion_image[:, :, :, :3], device=device, pbar=False)[0].to(dtype) + vae.to(offload_device) + + left_idx = r * infer_frames + right_idx = r * infer_frames + infer_frames + #cond_latents = COND[r] if pose_video else COND[0] * 0 + #cond_latents = cond_latents.to(dtype=self.param_dtype, device=self.device) + s2v_audio_input = s2v_audio_embeds[..., left_idx:right_idx] + input_motion_latents = ref_motion.clone() + + 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, s2v_ref_motion=input_motion_latents) + + latent = sample_scheduler.step( + noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0), + **scheduler_step_args)[0].squeeze(0) + + latents = torch.stack(latent) + #if not (drop_first_motion and r == 0): + # decode_latents = torch.cat([motion_latents, latents], dim=2) + #else: + decode_latents = torch.cat([s2v_ref_latent, latents], dim=2) + image = torch.stack(vae.decode(decode_latents), device=device) + image = image[:, :, -(infer_frames):] + #if (drop_first_motion and r == 0): + # image = image[:, :, 3:] + + 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(vae.device, vae.dtype) + + vae.to(device) + ref_motion_image = torch.stack(vae.encode(videos_last_frames, device=device, pbar=False)[0]) + vae.to(device) + framepack_out.append(image.cpu()) + + gen_video_samples = torch.cat(framepack_out, dim=1) + + 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.permute(1, 2, 3, 0)}, + #region normal inference else: noise_pred, self.cache_state = predict_with_cfg( diff --git a/nodes_model_loading.py b/nodes_model_loading.py index f053a8c..c5f9002 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -731,7 +731,7 @@ 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"} + params_to_keep = {"time_in", "patch_embedding", "time_", "modulation", "text_embedding", "adapter", "add", "ref_conv", "audio", "cond_encoder"} param_count = sum(1 for _ in transformer.named_parameters()) pbar = ProgressBar(param_count) cnt = 0 diff --git a/s2v/nodes.py b/s2v/nodes.py index 9bd8ab1..f854e16 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -51,12 +51,13 @@ class WanVideoAddAudioEmbeds: def INPUT_TYPES(s): return {"required": { "embeds": ("WANVIDIMAGE_EMBEDS",), - "audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",), "frames": ("INT", {"default": 81, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames to process"}), "audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for audio embeddings"}) }, "optional": { - "ref_latent": ("LATENT",) + "audio_encoder_output": ("AUDIO_ENCODER_OUTPUT",), + "ref_latent": ("LATENT",), + "pose_latent": ("LATENT",) } } @@ -66,38 +67,40 @@ class WanVideoAddAudioEmbeds: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, frames, audio_encoder_output, audio_scale, 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] + def add(self, embeds, frames, audio_encoder_output=None, audio_scale=1.0, ref_latent=None, pose_latent=None): + 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", audio_feat.shape) - input_fps = 50 - output_fps = 30 - bucket_fps = 16 + 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) + if input_fps != output_fps: + audio_feat = linear_interpolation(audio_feat, input_fps=input_fps, output_fps=output_fps) - self.video_rate = output_fps + self.video_rate = output_fps - audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps( - audio_feat, - fps=bucket_fps, - batch_frames=frames-1 - ) + audio_embed_bucket, num_repeat = self.get_audio_embed_bucket_fps( + audio_feat, + fps=bucket_fps, + batch_frames=frames-1 + ) - 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_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) - print("audio_embed_bucket", audio_embed_bucket.shape) + print("audio_embed_bucket", audio_embed_bucket.shape) new_entry = { - "audio_embed_bucket": audio_embed_bucket, - "num_repeat": num_repeat, + "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 } updated = dict(embeds) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 21f2a91..ce9cc2f 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2138,7 +2138,8 @@ class WanModel(torch.nn.Module): s2v_audio_input=None, s2v_ref_latent=None, s2v_audio_scale=1.0, - s2v_ref_motion=None + s2v_ref_motion=None, + s2v_pose=None ): r""" @@ -2243,6 +2244,12 @@ class WanModel(torch.nn.Module): for u in x ] + if s2v_pose is not None: + print("s2v_pose.shape:", s2v_pose.shape) + print("x[0].shape:", x[0].shape) + 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)] @@ -2371,8 +2378,8 @@ class WanModel(torch.nn.Module): 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) + #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: From 6faa24b7f2f90ddb2c0589ab4de6bf85357efeb1 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 28 Aug 2025 18:48:39 +0300 Subject: [PATCH 15/31] Implement Framepack long geneneration method RoPE handling is comfyanon's code --- nodes.py | 185 +- nodes_model_loading.py | 5 +- s2v/nodes.py | 52 +- ...anvideo2_2_S2V_context_window_testing.json | 1646 ++++---- ...anvideo2_2_S2V_framepack_pose_testing.json | 3448 +++++++++++++++++ wanvideo/modules/model.py | 49 +- wanvideo/schedulers/__init__.py | 3 - 7 files changed, 4497 insertions(+), 891 deletions(-) create mode 100644 s2v/wanvideo2_2_S2V_framepack_pose_testing.json diff --git a/nodes.py b/nodes.py index 50a570e..b7ad45b 100644 --- a/nodes.py +++ b/nodes.py @@ -1823,6 +1823,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) @@ -2223,14 +2225,19 @@ class WanVideoSampler: mtv_freqs = mtv_freqs.to(device, dtype) #region S2V - s2v_audio_input = s2v_ref_latent = None + 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"]].to(device, dtype) + #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: @@ -2241,10 +2248,10 @@ class WanVideoSampler: 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 = image_embeds.get("vae", None) - framepack = False + vae = s2v_audio_embeds.get("vae", None) # vid2vid noise_mask=original_image=None @@ -2525,7 +2532,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, s2v_audio_input=None, s2v_ref_motion=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) @@ -2670,7 +2677,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 @@ -2707,7 +2718,8 @@ class WanVideoSampler: "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_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 @@ -2862,7 +2874,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") @@ -3229,6 +3241,10 @@ class WanVideoSampler: 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) @@ -3246,7 +3262,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, s2v_audio_input=partial_s2v_audio_input) + 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 @@ -3671,67 +3687,124 @@ 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 - motion_frames = 5 - infer_frames = image_embeds["num_frames"] + #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 + + 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): if ref_motion_image is not None: - if ref_motion_image.shape[0] > 73: - ref_motion_image = ref_motion_image[-73:] - - if ref_motion_image.shape[0] < 73: - ref = torch.ones([73, ref_motion_image.shape[1], ref_motion_image.shape[2], 3]) * 0.5 - ref[-ref_motion_image.shape[0]:] = ref_motion_image - ref_motion_image = ref - vae.to(device) - ref_motion = vae.encode(ref_motion_image[:, :, :, :3], device=device, pbar=False)[0].to(dtype) + ref_motion = vae.encode(ref_motion_image.to(vae.dtype), device=device, pbar=False).to(dtype)[0] vae.to(offload_device) left_idx = r * infer_frames right_idx = r * infer_frames + infer_frames - #cond_latents = COND[r] if pose_video else COND[0] * 0 - #cond_latents = cond_latents.to(dtype=self.param_dtype, device=self.device) - s2v_audio_input = s2v_audio_embeds[..., left_idx:right_idx] - input_motion_latents = ref_motion.clone() - - 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, s2v_ref_motion=input_motion_latents) - latent = sample_scheduler.step( - noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0), - **scheduler_step_args)[0].squeeze(0) - - latents = torch.stack(latent) - #if not (drop_first_motion and r == 0): - # decode_latents = torch.cat([motion_latents, latents], dim=2) - #else: - decode_latents = torch.cat([s2v_ref_latent, latents], dim=2) - image = torch.stack(vae.decode(decode_latents), device=device) - image = image[:, :, -(infer_frames):] - #if (drop_first_motion and r == 0): - # image = image[:, :, 3:] + 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) - 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(vae.device, vae.dtype) - + if ref_motion_image is not None: + input_motion_latents = ref_motion.clone().unsqueeze(0) + else: + input_motion_latents = 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 + + vae.to(device) - ref_motion_image = torch.stack(vae.encode(videos_last_frames, device=device, pbar=False)[0]) - 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] + image = image.unsqueeze(0)[:, :, -infer_frames:] + if r == 0: + image = image[:, :, 3:] + framepack_out.append(image.cpu()) - gen_video_samples = torch.cat(framepack_out, dim=1) + 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) + 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"]: @@ -3741,7 +3814,7 @@ class WanVideoSampler: torch.cuda.reset_peak_memory_stats(device) except: pass - return {"video": gen_video_samples.permute(1, 2, 3, 0)}, + return {"video": gen_video_samples}, #region normal inference else: diff --git a/nodes_model_loading.py b/nodes_model_loading.py index c5f9002..4612e95 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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", "cond_encoder"} + 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 @@ -878,6 +879,7 @@ def patch_stand_in_lora(transformer, lora_sd, transformer_load_device, base_dtyp def add_lora_weights(patcher, lora, base_dtype, merge_loras=False): unianimate_sd = None + control_lora=False #spacepxl's control LoRA patch for l in lora: log.info(f"Loading LoRA: {l['name']} with strength: {l['strength']}") @@ -904,7 +906,6 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False): # Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys #if not any('img' in k for k in sd.keys()): # lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k} - control_lora=False if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd: control_lora = True #stand-in LoRA patch diff --git a/s2v/nodes.py b/s2v/nodes.py index f854e16..5c18cc6 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -46,47 +46,55 @@ def linear_interpolation(features, input_fps, output_fps, output_len=None): mode='linear') # [1, 512, output_len] return output_features.transpose(1, 2) # [1, output_len, 512] -class WanVideoAddAudioEmbeds: +class WanVideoAddS2VEmbeds: @classmethod def INPUT_TYPES(s): return {"required": { "embeds": ("WANVIDIMAGE_EMBEDS",), - "frames": ("INT", {"default": 81, "min": 1, "max": 100000, "step": 1, "tooltip": "Number of frames to process"}), - "audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.1, "tooltip": "Scale factor for audio embeddings"}) + "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",) + "pose_latent": ("LATENT",), + "vae": ("WANVAE",), + "enable_framepack": ("BOOLEAN", {"default": False, "tooltip": "Enable Framepack sampling loop, not compatible with context windows"}) } - } - RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",) - RETURN_NAMES = ("image_embeds",) + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "INT",) + RETURN_NAMES = ("image_embeds", "audio_frame_count") FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, frames, audio_encoder_output=None, audio_scale=1.0, ref_latent=None, pose_latent=None): + 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): 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", audio_feat.shape) - input_fps = 50 - output_fps = 30 - bucket_fps = 16 + 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=frames-1 + 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: @@ -94,6 +102,8 @@ class WanVideoAddAudioEmbeds: 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 = { @@ -101,11 +111,16 @@ class WanVideoAddAudioEmbeds: "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 + "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,) + 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 @@ -120,8 +135,7 @@ class WanVideoAddAudioEmbeds: 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 + 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, @@ -161,9 +175,9 @@ class WanVideoAddAudioEmbeds: NODE_CLASS_MAPPINGS = { - "WanVideoAddAudioEmbeds": WanVideoAddAudioEmbeds, + "WanVideoAddS2VEmbeds": WanVideoAddS2VEmbeds, } NODE_DISPLAY_NAME_MAPPINGS = { - "WanVideoAddAudioEmbeds": "WanVideo Add Audio Embeds", + "WanVideoAddS2VEmbeds": "WanVideo Add S2V Embeds", } \ No newline at end of file diff --git a/s2v/wanvideo2_2_S2V_context_window_testing.json b/s2v/wanvideo2_2_S2V_context_window_testing.json index 16e0a8e..d15e17a 100644 --- a/s2v/wanvideo2_2_S2V_context_window_testing.json +++ b/s2v/wanvideo2_2_S2V_context_window_testing.json @@ -1,8 +1,8 @@ { "id": "c6e410bc-5e2c-460b-ae81-c91b6094fbb1", "revision": 0, - "last_node_id": 100, - "last_link_id": 155, + "last_node_id": 102, + "last_link_id": 164, "nodes": [ { "id": 44, @@ -402,7 +402,7 @@ "type": "LATENT", "links": [ 92, - 128 + 157 ] } ], @@ -602,52 +602,6 @@ "color": "#223", "bgcolor": "#335" }, - { - "id": 83, - "type": "WanVideoContextOptions", - "pos": [ - 1290.159423828125, - -496.63787841796875 - ], - "size": [ - 275.783203125, - 202 - ], - "flags": {}, - "order": 8, - "mode": 0, - "inputs": [ - { - "name": "reference_latent", - "shape": 7, - "type": "LATENT", - "link": null - } - ], - "outputs": [ - { - "name": "context_options", - "type": "WANVIDCONTEXT", - "links": [ - 121 - ] - } - ], - "properties": { - "cnr_id": "ComfyUI-WanVideoWrapper", - "ver": "3c79851230c9ab042f52c8e176f349ecc3e51a64", - "Node name for S&R": "WanVideoContextOptions" - }, - "widgets_values": [ - "uniform_standard", - 73, - 4, - 16, - true, - false, - "linear" - ] - }, { "id": 77, "type": "InsertLatentToIndexed", @@ -758,7 +712,8 @@ "0, 0, 0", "center", 2, - "cpu" + "cpu", + "Output: 1 x 960 x 640 | 7.03MB" ] }, { @@ -773,7 +728,7 @@ 202 ], "flags": {}, - "order": 9, + "order": 8, "mode": 0, "inputs": [], "outputs": [ @@ -845,7 +800,7 @@ "links": null }, { - "label": "601 count", + "label": "201 count", "name": "count", "type": "INT", "links": null @@ -870,7 +825,7 @@ 342 ], "flags": {}, - "order": 10, + "order": 9, "mode": 0, "inputs": [ { @@ -917,6 +872,503 @@ "color": "#223", "bgcolor": "#335" }, + { + "id": 94, + "type": "VHS_LoadAudio", + "pos": [ + -59.62910079956055, + 496.306396484375 + ], + "size": [ + 415.6340026855469, + 126 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "audio", + "type": "AUDIO", + "links": [ + 148, + 153, + 154 + ] + }, + { + "name": "duration", + "type": "FLOAT", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "8e4d79471bf1952154768e8435a9300077b534fa", + "Node name for S&R": "VHS_LoadAudio" + }, + "widgets_values": { + "audio_file": "input/weightoftheworld2.mp4", + "seek_seconds": 0, + "duration": 0 + } + }, + { + "id": 66, + "type": "LoadAudio", + "pos": [ + -40.06507873535156, + 719.501953125 + ], + "size": [ + 361.2844543457031, + 155.81912231445312 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "AUDIO", + "type": "AUDIO", + "links": [] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "LoadAudio" + }, + "widgets_values": [ + "NieR_ Automata - _Weight of the World_ ENG VER. by Lizz Robinett [CyOSTbel3AM].mp3", + null, + null + ] + }, + { + "id": 81, + "type": "MelBandRoFormerModelLoader", + "pos": [ + 563.452392578125, + 647.8037109375 + ], + "size": [ + 316.2164001464844, + 58 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "model", + "type": "MELROFORMERMODEL", + "links": [ + 106 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-MelBandRoFormer", + "ver": "b40e263224778ec417114d91d8b3b39934e30de5", + "Node name for S&R": "MelBandRoFormerModelLoader" + }, + "widgets_values": [ + "MelBandRoFormer\\MelBandRoformer_fp16.safetensors" + ] + }, + { + "id": 82, + "type": "MelBandRoFormerSampler", + "pos": [ + 555.0003662109375, + 770.088623046875 + ], + "size": [ + 222.73397827148438, + 46 + ], + "flags": {}, + "order": 19, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MELROFORMERMODEL", + "link": 106 + }, + { + "name": "audio", + "type": "AUDIO", + "link": 154 + } + ], + "outputs": [ + { + "name": "vocals", + "type": "AUDIO", + "links": [ + 149 + ] + }, + { + "name": "instruments", + "type": "AUDIO", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-MelBandRoFormer", + "ver": "b40e263224778ec417114d91d8b3b39934e30de5", + "Node name for S&R": "MelBandRoFormerSampler" + }, + "widgets_values": [] + }, + { + "id": 98, + "type": "NormalizeAudioLoudness", + "pos": [ + 834.3135986328125, + 760.9041137695312 + ], + "size": [ + 270, + 58 + ], + "flags": {}, + "order": 24, + "mode": 0, + "inputs": [ + { + "name": "audio", + "type": "AUDIO", + "link": 149 + } + ], + "outputs": [ + { + "name": "audio", + "type": "AUDIO", + "links": [ + 150 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "90c3bbb6c2e4ff5e05305e765d007d5e58428ce4", + "Node name for S&R": "NormalizeAudioLoudness" + }, + "widgets_values": [ + -23 + ] + }, + { + "id": 95, + "type": "DownloadAndLoadGIMMVFIModel", + "pos": [ + 3262.200927734375, + -713.7930908203125 + ], + "size": [ + 339.4301452636719, + 106 + ], + "flags": {}, + "order": 13, + "mode": 2, + "inputs": [], + "outputs": [ + { + "name": "gimmvfi_model", + "type": "GIMMVIF_MODEL", + "links": [ + 144 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-GIMM-VFI", + "ver": "4c9a3123762af85e7c796e41737da0b70c75d72d", + "Node name for S&R": "DownloadAndLoadGIMMVFIModel" + }, + "widgets_values": [ + "gimmvfi_r_arb_lpips_fp32.safetensors", + "fp16", + false + ] + }, + { + "id": 80, + "type": "VHS_SplitImages", + "pos": [ + 2397.607666015625, + -214.64483642578125 + ], + "size": [ + 210, + 118 + ], + "flags": {}, + "order": 33, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 100 + } + ], + "outputs": [ + { + "name": "IMAGE_A", + "type": "IMAGE", + "links": null + }, + { + "name": "A_count", + "type": "INT", + "links": null + }, + { + "name": "IMAGE_B", + "type": "IMAGE", + "links": [ + 145, + 155 + ] + }, + { + "name": "B_count", + "type": "INT", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "8e4d79471bf1952154768e8435a9300077b534fa", + "Node name for S&R": "VHS_SplitImages" + }, + "widgets_values": { + "split_index": 3 + } + }, + { + "id": 64, + "type": "AudioEncoderEncode", + "pos": [ + 1198.93212890625, + 778.826904296875 + ], + "size": [ + 285.087890625, + 46 + ], + "flags": {}, + "order": 26, + "mode": 0, + "inputs": [ + { + "name": "audio_encoder", + "type": "AUDIO_ENCODER", + "link": 69 + }, + { + "name": "audio", + "type": "AUDIO", + "link": 150 + } + ], + "outputs": [ + { + "name": "AUDIO_ENCODER_OUTPUT", + "type": "AUDIO_ENCODER_OUTPUT", + "links": [ + 156 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "AudioEncoderEncode" + }, + "widgets_values": [] + }, + { + "id": 65, + "type": "AudioEncoderLoader", + "pos": [ + 1136.3577880859375, + 641.1036376953125 + ], + "size": [ + 346.756103515625, + 58 + ], + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "AUDIO_ENCODER", + "type": "AUDIO_ENCODER", + "links": [ + 69 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "AudioEncoderLoader" + }, + "widgets_values": [ + "wav2vec_xlsr_53_english_fp32.safetensors" + ] + }, + { + "id": 69, + "type": "PreviewAny", + "pos": [ + 2732.885009765625, + -1020.7718505859375 + ], + "size": [ + 490.69390869140625, + 355.7907409667969 + ], + "flags": {}, + "order": 29, + "mode": 0, + "inputs": [ + { + "name": "source", + "type": "*", + "link": 159 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "PreviewAny" + }, + "widgets_values": [] + }, + { + "id": 37, + "type": "WanVideoEmptyEmbeds", + "pos": [ + 1442.7791748046875, + -1070.2607421875 + ], + "size": [ + 315, + 126 + ], + "flags": {}, + "order": 21, + "mode": 0, + "inputs": [ + { + "name": "control_embeds", + "shape": 7, + "type": "WANVIDIMAGE_EMBEDS", + "link": null + }, + { + "name": "extra_latents", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "width", + "type": "INT", + "widget": { + "name": "width" + }, + "link": 85 + }, + { + "name": "height", + "type": "INT", + "widget": { + "name": "height" + }, + "link": 86 + }, + { + "name": "num_frames", + "type": "INT", + "widget": { + "name": "num_frames" + }, + "link": 79 + } + ], + "outputs": [ + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "links": [ + 158 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoEmptyEmbeds" + }, + "widgets_values": [ + 832, + 480, + 201 + ] + }, + { + "id": 71, + "type": "PrimitiveNode", + "pos": [ + 1388.9095458984375, + -791.9364013671875 + ], + "size": [ + 210, + 82 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "widget": { + "name": "num_frames" + }, + "links": [ + 79, + 161 + ] + } + ], + "title": "num_frames", + "properties": { + "Run widget replace on values": false + }, + "widgets_values": [ + 201, + "fixed" + ] + }, { "id": 27, "type": "WanVideoSampler", @@ -929,7 +1381,7 @@ 900.6666870117188 ], "flags": {}, - "order": 29, + "order": 28, "mode": 0, "inputs": [ { @@ -940,7 +1392,7 @@ { "name": "image_embeds", "type": "WANVIDIMAGE_EMBEDS", - "link": 91 + "link": 162 }, { "name": "text_embeds", @@ -1055,13 +1507,13 @@ "Node name for S&R": "WanVideoSampler" }, "widgets_values": [ - 6, + 4, 1, 4, 45, "fixed", true, - "lcm", + "dpm++_sde", 0, 1, false, @@ -1072,205 +1524,314 @@ ] }, { - "id": 94, - "type": "VHS_LoadAudio", + "id": 101, + "type": "WanVideoAddS2VEmbeds", "pos": [ - -59.62910079956055, - 496.306396484375 + 2000.4896240234375, + -871.2449340820312 ], "size": [ - 415.6340026855469, - 126 + 327.17578125, + 234 ], "flags": {}, - "order": 11, + "order": 27, "mode": 0, - "inputs": [], + "inputs": [ + { + "name": "embeds", + "type": "WANVIDIMAGE_EMBEDS", + "link": 158 + }, + { + "name": "audio_encoder_output", + "shape": 7, + "type": "AUDIO_ENCODER_OUTPUT", + "link": 156 + }, + { + "name": "ref_latent", + "shape": 7, + "type": "LATENT", + "link": 157 + }, + { + "name": "pose_latent", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "WANVAE", + "link": null + }, + { + "name": "frame_window_size", + "type": "INT", + "widget": { + "name": "frame_window_size" + }, + "link": 161 + } + ], "outputs": [ { - "name": "audio", - "type": "AUDIO", + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", "links": [ - 148, - 153, - 154 + 162 ] }, { - "name": "duration", - "type": "FLOAT", + "name": "audio_frame_count", + "type": "INT", + "links": [ + 159 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "a5621b87391013155b4f688fbe01dba10a8104aa", + "Node name for S&R": "WanVideoAddS2VEmbeds" + }, + "widgets_values": [ + 201, + 1, + 0, + 1, + false + ] + }, + { + "id": 83, + "type": "WanVideoContextOptions", + "pos": [ + 1290.159423828125, + -496.63787841796875 + ], + "size": [ + 275.783203125, + 202 + ], + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [ + { + "name": "reference_latent", + "shape": 7, + "type": "LATENT", + "link": null + } + ], + "outputs": [ + { + "name": "context_options", + "type": "WANVIDCONTEXT", + "links": [ + 121 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "3c79851230c9ab042f52c8e176f349ecc3e51a64", + "Node name for S&R": "WanVideoContextOptions" + }, + "widgets_values": [ + "uniform_standard", + 81, + 4, + 16, + true, + false, + "linear" + ] + }, + { + "id": 96, + "type": "GIMMVFI_interpolate", + "pos": [ + 3209.548583984375, + -509.1043701171875 + ], + "size": [ + 270, + 174 + ], + "flags": {}, + "order": 34, + "mode": 2, + "inputs": [ + { + "name": "gimmvfi_model", + "type": "GIMMVIF_MODEL", + "link": 144 + }, + { + "name": "images", + "type": "IMAGE", + "link": 145 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 163 + ] + }, + { + "name": "flow_tensors", + "type": "IMAGE", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-GIMM-VFI", + "ver": "4c9a3123762af85e7c796e41737da0b70c75d72d", + "Node name for S&R": "GIMMVFI_interpolate" + }, + "widgets_values": [ + 1, + 3, + 0, + "fixed", + false + ] + }, + { + "id": 30, + "type": "VHS_VideoCombine", + "pos": [ + 3815.402099609375, + -864.041748046875 + ], + "size": [ + 940.8292846679688, + 334 + ], + "flags": {}, + "order": 37, + "mode": 2, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 164 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": 153 + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "330bce6c3c0d47ebdedcc0348d9ab355707b7523", + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 24, + "loop_count": 0, + "filename_prefix": "WanVideo2_2_S2V", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "trim_to_audio": false, + "pingpong": false, + "save_output": false, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "WanVideo2_2_S2V_00013-audio.mp4", + "subfolder": "", + "type": "temp", + "format": "video/h264-mp4", + "frame_rate": 32, + "workflow": "WanVideo2_2_S2V_00013.png", + "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo2_2_S2V_00013-audio.mp4" + } + } + } + }, + { + "id": 102, + "type": "VHS_SelectEveryNthImage", + "pos": [ + 3513.563720703125, + -508.5723876953125 + ], + "size": [ + 266.349609375, + 102 + ], + "flags": {}, + "order": 36, + "mode": 2, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 163 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 164 + ] + }, + { + "name": "count", + "type": "INT", "links": null } ], "properties": { "cnr_id": "comfyui-videohelpersuite", "ver": "8e4d79471bf1952154768e8435a9300077b534fa", - "Node name for S&R": "VHS_LoadAudio" + "Node name for S&R": "VHS_SelectEveryNthImage" }, "widgets_values": { - "audio_file": "input/weightoftheworld2.mp4", - "seek_seconds": 0, - "duration": 0 + "select_every_nth": 2, + "skip_first_images": 0 } }, - { - "id": 66, - "type": "LoadAudio", - "pos": [ - -40.06507873535156, - 719.501953125 - ], - "size": [ - 361.2844543457031, - 155.81912231445312 - ], - "flags": {}, - "order": 12, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "AUDIO", - "type": "AUDIO", - "links": [] - } - ], - "properties": { - "cnr_id": "comfy-core", - "ver": "0.3.52", - "Node name for S&R": "LoadAudio" - }, - "widgets_values": [ - "NieR_ Automata - _Weight of the World_ ENG VER. by Lizz Robinett [CyOSTbel3AM].mp3", - null, - null - ] - }, - { - "id": 81, - "type": "MelBandRoFormerModelLoader", - "pos": [ - 563.452392578125, - 647.8037109375 - ], - "size": [ - 316.2164001464844, - 58 - ], - "flags": {}, - "order": 13, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "model", - "type": "MELROFORMERMODEL", - "links": [ - 106 - ] - } - ], - "properties": { - "cnr_id": "ComfyUI-MelBandRoFormer", - "ver": "b40e263224778ec417114d91d8b3b39934e30de5", - "Node name for S&R": "MelBandRoFormerModelLoader" - }, - "widgets_values": [ - "MelBandRoFormer\\MelBandRoformer_fp16.safetensors" - ] - }, - { - "id": 82, - "type": "MelBandRoFormerSampler", - "pos": [ - 555.0003662109375, - 770.088623046875 - ], - "size": [ - 222.73397827148438, - 46 - ], - "flags": {}, - "order": 19, - "mode": 0, - "inputs": [ - { - "name": "model", - "type": "MELROFORMERMODEL", - "link": 106 - }, - { - "name": "audio", - "type": "AUDIO", - "link": 154 - } - ], - "outputs": [ - { - "name": "vocals", - "type": "AUDIO", - "links": [ - 149 - ] - }, - { - "name": "instruments", - "type": "AUDIO", - "links": null - } - ], - "properties": { - "cnr_id": "ComfyUI-MelBandRoFormer", - "ver": "b40e263224778ec417114d91d8b3b39934e30de5", - "Node name for S&R": "MelBandRoFormerSampler" - }, - "widgets_values": [] - }, - { - "id": 98, - "type": "NormalizeAudioLoudness", - "pos": [ - 834.3135986328125, - 760.9041137695312 - ], - "size": [ - 270, - 58 - ], - "flags": {}, - "order": 24, - "mode": 0, - "inputs": [ - { - "name": "audio", - "type": "AUDIO", - "link": 149 - } - ], - "outputs": [ - { - "name": "audio", - "type": "AUDIO", - "links": [ - 150 - ] - } - ], - "properties": { - "cnr_id": "ComfyUI-WanVideoWrapper", - "ver": "90c3bbb6c2e4ff5e05305e765d007d5e58428ce4", - "Node name for S&R": "NormalizeAudioLoudness" - }, - "widgets_values": [ - -23 - ] - }, { "id": 97, "type": "VHS_VideoCombine", "pos": [ - 3319.513671875, - 193.33419799804688 + 2914.45703125, + -176.39088439941406 ], "size": [ 940.8292846679688, @@ -1331,509 +1892,16 @@ "hidden": false, "paused": false, "params": { - "filename": "WanVideo2_2_S2V_00012-audio.mp4", + "filename": "WanVideo2_2_S2V_00015-audio.mp4", "subfolder": "", "type": "temp", "format": "video/h264-mp4", "frame_rate": 16, - "workflow": "WanVideo2_2_S2V_00012.png", - "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo2_2_S2V_00012-audio.mp4" + "workflow": "WanVideo2_2_S2V_00015.png", + "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo2_2_S2V_00015-audio.mp4" } } } - }, - { - "id": 30, - "type": "VHS_VideoCombine", - "pos": [ - 3649.713623046875, - -868.4019165039062 - ], - "size": [ - 940.8292846679688, - 961.8861694335938 - ], - "flags": {}, - "order": 36, - "mode": 2, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 146 - }, - { - "name": "audio", - "shape": 7, - "type": "AUDIO", - "link": 153 - }, - { - "name": "meta_batch", - "shape": 7, - "type": "VHS_BatchManager", - "link": null - }, - { - "name": "vae", - "shape": 7, - "type": "VAE", - "link": null - } - ], - "outputs": [ - { - "name": "Filenames", - "type": "VHS_FILENAMES", - "links": null - } - ], - "properties": { - "cnr_id": "comfyui-videohelpersuite", - "ver": "330bce6c3c0d47ebdedcc0348d9ab355707b7523", - "Node name for S&R": "VHS_VideoCombine" - }, - "widgets_values": { - "frame_rate": 32, - "loop_count": 0, - "filename_prefix": "WanVideo2_2_S2V", - "format": "video/h264-mp4", - "pix_fmt": "yuv420p", - "crf": 19, - "save_metadata": true, - "trim_to_audio": false, - "pingpong": false, - "save_output": false, - "videopreview": { - "hidden": false, - "paused": false, - "params": { - "filename": "WanVideo2_2_S2V_00013-audio.mp4", - "subfolder": "", - "type": "temp", - "format": "video/h264-mp4", - "frame_rate": 32, - "workflow": "WanVideo2_2_S2V_00013.png", - "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo2_2_S2V_00013-audio.mp4" - } - } - } - }, - { - "id": 96, - "type": "GIMMVFI_interpolate", - "pos": [ - 3256.453369140625, - -512.8270263671875 - ], - "size": [ - 270, - 174 - ], - "flags": {}, - "order": 34, - "mode": 2, - "inputs": [ - { - "name": "gimmvfi_model", - "type": "GIMMVIF_MODEL", - "link": 144 - }, - { - "name": "images", - "type": "IMAGE", - "link": 145 - } - ], - "outputs": [ - { - "name": "images", - "type": "IMAGE", - "links": [ - 146 - ] - }, - { - "name": "flow_tensors", - "type": "IMAGE", - "links": null - } - ], - "properties": { - "cnr_id": "ComfyUI-GIMM-VFI", - "ver": "4c9a3123762af85e7c796e41737da0b70c75d72d", - "Node name for S&R": "GIMMVFI_interpolate" - }, - "widgets_values": [ - 1, - 2, - 0, - "fixed", - false - ] - }, - { - "id": 95, - "type": "DownloadAndLoadGIMMVFIModel", - "pos": [ - 3262.200927734375, - -713.7930908203125 - ], - "size": [ - 339.4301452636719, - 106 - ], - "flags": {}, - "order": 14, - "mode": 2, - "inputs": [], - "outputs": [ - { - "name": "gimmvfi_model", - "type": "GIMMVIF_MODEL", - "links": [ - 144 - ] - } - ], - "properties": { - "cnr_id": "ComfyUI-GIMM-VFI", - "ver": "4c9a3123762af85e7c796e41737da0b70c75d72d", - "Node name for S&R": "DownloadAndLoadGIMMVFIModel" - }, - "widgets_values": [ - "gimmvfi_r_arb_lpips_fp32.safetensors", - "fp16", - false - ] - }, - { - "id": 80, - "type": "VHS_SplitImages", - "pos": [ - 2397.607666015625, - -214.64483642578125 - ], - "size": [ - 210, - 118 - ], - "flags": {}, - "order": 33, - "mode": 0, - "inputs": [ - { - "name": "images", - "type": "IMAGE", - "link": 100 - } - ], - "outputs": [ - { - "name": "IMAGE_A", - "type": "IMAGE", - "links": null - }, - { - "name": "A_count", - "type": "INT", - "links": null - }, - { - "name": "IMAGE_B", - "type": "IMAGE", - "links": [ - 145, - 155 - ] - }, - { - "name": "B_count", - "type": "INT", - "links": null - } - ], - "properties": { - "cnr_id": "comfyui-videohelpersuite", - "ver": "8e4d79471bf1952154768e8435a9300077b534fa", - "Node name for S&R": "VHS_SplitImages" - }, - "widgets_values": { - "split_index": 3 - } - }, - { - "id": 63, - "type": "WanVideoAddAudioEmbeds", - "pos": [ - 2055.220703125, - -847.4767456054688 - ], - "size": [ - 254.5941619873047, - 122 - ], - "flags": {}, - "order": 27, - "mode": 0, - "inputs": [ - { - "name": "embeds", - "type": "WANVIDIMAGE_EMBEDS", - "link": 66 - }, - { - "name": "audio_encoder_output", - "type": "AUDIO_ENCODER_OUTPUT", - "link": 68 - }, - { - "name": "ref_latent", - "shape": 7, - "type": "LATENT", - "link": 128 - }, - { - "name": "frames", - "type": "INT", - "widget": { - "name": "frames" - }, - "link": 96 - } - ], - "outputs": [ - { - "name": "image_embeds", - "type": "WANVIDIMAGE_EMBEDS", - "links": [ - 75, - 91 - ] - } - ], - "properties": { - "cnr_id": "ComfyUI-WanVideoWrapper", - "ver": "a1ca0985ec120ff97e34676a64de19c99767bbd4", - "Node name for S&R": "WanVideoAddAudioEmbeds" - }, - "widgets_values": [ - 601, - 1 - ] - }, - { - "id": 69, - "type": "PreviewAny", - "pos": [ - 2441.904052734375, - -1017.1683349609375 - ], - "size": [ - 490.69390869140625, - 355.7907409667969 - ], - "flags": {}, - "order": 28, - "mode": 0, - "inputs": [ - { - "name": "source", - "type": "*", - "link": 75 - } - ], - "outputs": [], - "properties": { - "cnr_id": "comfy-core", - "ver": "0.3.52", - "Node name for S&R": "PreviewAny" - }, - "widgets_values": [] - }, - { - "id": 37, - "type": "WanVideoEmptyEmbeds", - "pos": [ - 1686.0142822265625, - -1071.16162109375 - ], - "size": [ - 315, - 126 - ], - "flags": {}, - "order": 21, - "mode": 0, - "inputs": [ - { - "name": "control_embeds", - "shape": 7, - "type": "WANVIDIMAGE_EMBEDS", - "link": null - }, - { - "name": "extra_latents", - "shape": 7, - "type": "LATENT", - "link": null - }, - { - "name": "width", - "type": "INT", - "widget": { - "name": "width" - }, - "link": 85 - }, - { - "name": "height", - "type": "INT", - "widget": { - "name": "height" - }, - "link": 86 - }, - { - "name": "num_frames", - "type": "INT", - "widget": { - "name": "num_frames" - }, - "link": 79 - } - ], - "outputs": [ - { - "name": "image_embeds", - "type": "WANVIDIMAGE_EMBEDS", - "links": [ - 66 - ] - } - ], - "properties": { - "cnr_id": "ComfyUI-WanVideoWrapper", - "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", - "Node name for S&R": "WanVideoEmptyEmbeds" - }, - "widgets_values": [ - 832, - 480, - 601 - ] - }, - { - "id": 71, - "type": "PrimitiveNode", - "pos": [ - 1682.593505859375, - -863.1051635742188 - ], - "size": [ - 210, - 82 - ], - "flags": {}, - "order": 15, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "INT", - "type": "INT", - "widget": { - "name": "num_frames" - }, - "links": [ - 79, - 96 - ] - } - ], - "title": "num_frames", - "properties": { - "Run widget replace on values": false - }, - "widgets_values": [ - 601, - "fixed" - ] - }, - { - "id": 64, - "type": "AudioEncoderEncode", - "pos": [ - 1198.93212890625, - 778.826904296875 - ], - "size": [ - 285.087890625, - 46 - ], - "flags": {}, - "order": 26, - "mode": 0, - "inputs": [ - { - "name": "audio_encoder", - "type": "AUDIO_ENCODER", - "link": 69 - }, - { - "name": "audio", - "type": "AUDIO", - "link": 150 - } - ], - "outputs": [ - { - "name": "AUDIO_ENCODER_OUTPUT", - "type": "AUDIO_ENCODER_OUTPUT", - "links": [ - 68 - ] - } - ], - "properties": { - "cnr_id": "comfy-core", - "ver": "0.3.52", - "Node name for S&R": "AudioEncoderEncode" - }, - "widgets_values": [] - }, - { - "id": 65, - "type": "AudioEncoderLoader", - "pos": [ - 1136.3577880859375, - 641.1036376953125 - ], - "size": [ - 346.756103515625, - 58 - ], - "flags": {}, - "order": 16, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "AUDIO_ENCODER", - "type": "AUDIO_ENCODER", - "links": [ - 69 - ] - } - ], - "properties": { - "cnr_id": "comfy-core", - "ver": "0.3.52", - "Node name for S&R": "AudioEncoderLoader" - }, - "widgets_values": [ - "wav2vec_xlsr_53_english_fp32.safetensors" - ] } ], "links": [ @@ -1893,22 +1961,6 @@ 0, "*" ], - [ - 66, - 37, - 0, - 63, - 0, - "WANVIDIMAGE_EMBEDS" - ], - [ - 68, - 64, - 0, - 63, - 1, - "AUDIO_ENCODER_OUTPUT" - ], [ 69, 65, @@ -1925,14 +1977,6 @@ 2, "WANVIDEOTEXTEMBEDS" ], - [ - 75, - 63, - 0, - 69, - 0, - "*" - ], [ 77, 28, @@ -1989,14 +2033,6 @@ 3, "INT" ], - [ - 91, - 63, - 0, - 27, - 1, - "WANVIDIMAGE_EMBEDS" - ], [ 92, 72, @@ -2013,14 +2049,6 @@ 1, "LATENT" ], - [ - 96, - 71, - 0, - 63, - 3, - "INT" - ], [ 100, 70, @@ -2053,14 +2081,6 @@ 5, "WANVIDCONTEXT" ], - [ - 128, - 72, - 0, - 63, - 2, - "LATENT" - ], [ 129, 35, @@ -2085,14 +2105,6 @@ 1, "IMAGE" ], - [ - 146, - 96, - 0, - 30, - 0, - "IMAGE" - ], [ 148, 94, @@ -2140,16 +2152,80 @@ 97, 0, "IMAGE" + ], + [ + 156, + 64, + 0, + 101, + 1, + "AUDIO_ENCODER_OUTPUT" + ], + [ + 157, + 72, + 0, + 101, + 2, + "LATENT" + ], + [ + 158, + 37, + 0, + 101, + 0, + "WANVIDIMAGE_EMBEDS" + ], + [ + 159, + 101, + 1, + 69, + 0, + "*" + ], + [ + 161, + 71, + 0, + 101, + 5, + "INT" + ], + [ + 162, + 101, + 0, + 27, + 1, + "WANVIDIMAGE_EMBEDS" + ], + [ + 163, + 96, + 0, + 102, + 0, + "IMAGE" + ], + [ + 164, + 102, + 0, + 30, + 0, + "IMAGE" ] ], "groups": [], "config": {}, "extra": { "ds": { - "scale": 0.505447028499326, + "scale": 0.5116450399201969, "offset": [ - 740.042928663884, - 1057.007052457307 + 572.0674176111402, + 1097.812069895393 ] }, "frontendVersion": "1.26.6", diff --git a/s2v/wanvideo2_2_S2V_framepack_pose_testing.json b/s2v/wanvideo2_2_S2V_framepack_pose_testing.json new file mode 100644 index 0000000..4008de1 --- /dev/null +++ b/s2v/wanvideo2_2_S2V_framepack_pose_testing.json @@ -0,0 +1,3448 @@ +{ + "id": "c6e410bc-5e2c-460b-ae81-c91b6094fbb1", + "revision": 0, + "last_node_id": 146, + "last_link_id": 244, + "nodes": [ + { + "id": 28, + "type": "WanVideoDecode", + "pos": [ + 1994.2247314453125, + -394.9518737792969 + ], + "size": [ + 315, + 198 + ], + "flags": {}, + "order": 53, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "WANVAE", + "link": 197 + }, + { + "name": "samples", + "type": "LATENT", + "link": 105 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 77 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoDecode" + }, + "widgets_values": [ + false, + 272, + 272, + 144, + 128, + "default" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 35, + "type": "WanVideoTorchCompileSettings", + "pos": [ + -1117.06494140625, + -1156.921875 + ], + "size": [ + 390.5999755859375, + 202 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "torch_compile_args", + "type": "WANCOMPILEARGS", + "slot_index": 0, + "links": [ + 129 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoTorchCompileSettings" + }, + "widgets_values": [ + "inductor", + false, + "default", + false, + 64, + true, + 128 + ] + }, + { + "id": 44, + "type": "Note", + "pos": [ + -1106.779052734375, + -1309.5294189453125 + ], + "size": [ + 303.0501403808594, + 88 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "If you have Triton installed, connect this for ~20-30% speed increase and reduced peak VRAM use" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 36, + "type": "Note", + "pos": [ + -559.8602294921875, + -1394.335205078125 + ], + "size": [ + 374.3061828613281, + 171.9547576904297 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "fp_16_fast enables \"Full FP16 Accmumulation in FP16 GEMMs\" feature available in the very latest pytorch nightly, this is around 20% speed boost. \n\nSageattn if you have it installed can be used for almost double inference speed at higher resolutions\n\nRadial attention is even faster but has worst quality, it should be used along with Set Radial Attention node to control which steps/blocks it's applied on to balance quality and speed." + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 22, + "type": "WanVideoModelLoader", + "pos": [ + -593.1958618164062, + -1146.4970703125 + ], + "size": [ + 477.4410095214844, + 314 + ], + "flags": {}, + "order": 29, + "mode": 0, + "inputs": [ + { + "name": "compile_args", + "shape": 7, + "type": "WANCOMPILEARGS", + "link": 129 + }, + { + "name": "block_swap_args", + "shape": 7, + "type": "BLOCKSWAPARGS", + "link": null + }, + { + "name": "lora", + "shape": 7, + "type": "WANVIDLORA", + "link": null + }, + { + "name": "vram_management_args", + "shape": 7, + "type": "VRAM_MANAGEMENTARGS", + "link": null + }, + { + "name": "extra_model", + "shape": 7, + "type": "VACEPATH", + "link": null + }, + { + "name": "fantasytalking_model", + "shape": 7, + "type": "FANTASYTALKINGMODEL", + "link": null + }, + { + "name": "multitalk_model", + "shape": 7, + "type": "MULTITALKMODEL", + "link": null + }, + { + "name": "fantasyportrait_model", + "shape": 7, + "type": "FANTASYPORTRAITMODEL", + "link": null + }, + { + "name": "vace_model", + "shape": 7, + "type": "VACEPATH", + "link": null + } + ], + "outputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "slot_index": 0, + "links": [ + 61 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoModelLoader" + }, + "widgets_values": [ + "WanVideo\\S2V\\Wan2_2-S2V-14B_fp8_e4m3fn_scaled_KJ.safetensors", + "fp16_fast", + "fp8_e4m3fn_scaled", + "offload_device", + "sageattn" + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 61, + "type": "MarkdownNote", + "pos": [ + -83.98224639892578, + -1481.6065673828125 + ], + "size": [ + 688.705078125, + 156.6822052001953 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "Models:\n\n[https://huggingface.co/Kijai/WanVideo_comfy_fp8_scaled](https://huggingface.co/Kijai/WanVideo_comfy_fp8_scaled)\n\nIf you want to use torch compile on GPUs prior to 4000 series:\n\nLoRA:\n\n[https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Lightx2v/lightx2v_T2V_14B_cfg_step_distill_v2_lora_rank64_bf16.safetensors](https://huggingface.co/Kijai/WanVideo_comfy/blob/main/Lightx2v/lightx2v_T2V_14B_cfg_step_distill_v2_lora_rank64_bf16.safetensors)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 60, + "type": "WanVideoLoraSelectMulti", + "pos": [ + 17.099323272705078, + -1224.06787109375 + ], + "size": [ + 617.9208374023438, + 342 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "prev_lora", + "shape": 7, + "type": "WANVIDLORA", + "link": null + }, + { + "name": "blocks", + "shape": 7, + "type": "SELECTEDBLOCKS", + "link": null + } + ], + "outputs": [ + { + "name": "lora", + "type": "WANVIDLORA", + "links": [ + 64 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoLoraSelectMulti" + }, + "widgets_values": [ + "WanVideo\\Lightx2v\\lightx2v_T2V_14B_cfg_step_distill_v2_lora_rank64_bf16_.safetensors", + 1.2, + "none", + 1, + "none", + 1, + "none", + 1, + "none", + 1, + false, + false + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 58, + "type": "WanVideoSetLoRAs", + "pos": [ + 36.42268371582031, + -805.3877563476562 + ], + "size": [ + 174.53378295898438, + 46 + ], + "flags": {}, + "order": 36, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 61 + }, + { + "name": "lora", + "shape": 7, + "type": "WANVIDLORA", + "link": 64 + } + ], + "outputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "links": [ + 62 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoSetLoRAs" + }, + "widgets_values": [], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 39, + "type": "WanVideoBlockSwap", + "pos": [ + -434.31292724609375, + -777.4451293945312 + ], + "size": [ + 315, + 202 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "block_swap_args", + "type": "BLOCKSWAPARGS", + "slot_index": 0, + "links": [ + 58 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoBlockSwap" + }, + "widgets_values": [ + 32, + false, + false, + true, + 0, + 1, + false + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 42, + "type": "Note", + "pos": [ + -765.3607788085938, + -775.6983032226562 + ], + "size": [ + 312.98052978515625, + 92.32489013671875 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "Adjust the blocks to swap based on your VRAM, this is a tradeoff between speed and memory usage." + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 56, + "type": "WanVideoSetBlockSwap", + "pos": [ + 250.55487060546875, + -803.0498657226562 + ], + "size": [ + 201.76815795898438, + 46 + ], + "flags": {}, + "order": 42, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 62 + }, + { + "name": "block_swap_args", + "shape": 7, + "type": "BLOCKSWAPARGS", + "link": 58 + } + ], + "outputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "links": [ + 60 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoSetBlockSwap" + }, + "widgets_values": [], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 38, + "type": "WanVideoVAELoader", + "pos": [ + 20.023881912231445, + -646.4891357421875 + ], + "size": [ + 315, + 82 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "compile_args", + "shape": 7, + "type": "WANCOMPILEARGS", + "link": null + } + ], + "outputs": [ + { + "name": "vae", + "type": "WANVAE", + "slot_index": 0, + "links": [ + 196 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoVAELoader" + }, + "widgets_values": [ + "wanvideo\\Wan2_1_VAE_bf16.safetensors", + "bf16" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 119, + "type": "SetNode", + "pos": [ + 391.3757019042969, + -615.3194580078125 + ], + "size": [ + 210, + 60 + ], + "flags": { + "collapsed": true + }, + "order": 30, + "mode": 0, + "inputs": [ + { + "name": "WANVAE", + "type": "WANVAE", + "link": 196 + } + ], + "outputs": [ + { + "name": "*", + "type": "*", + "links": null + } + ], + "title": "Set_VAE", + "properties": { + "previousName": "VAE" + }, + "widgets_values": [ + "VAE" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 120, + "type": "GetNode", + "pos": [ + 2002.47412109375, + -472.52252197265625 + ], + "size": [ + 210, + 60 + ], + "flags": { + "collapsed": true + }, + "order": 8, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "WANVAE", + "type": "WANVAE", + "links": [ + 197 + ] + } + ], + "title": "Get_VAE", + "properties": {}, + "widgets_values": [ + "VAE" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 109, + "type": "WanVideoEncode", + "pos": [ + 226.22377014160156, + 1413.26171875 + ], + "size": [ + 270, + 242 + ], + "flags": {}, + "order": 48, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "WANVAE", + "link": 199 + }, + { + "name": "image", + "type": "IMAGE", + "link": 174 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "links": [ + 192 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "63d4b6aadaae543f96d101788122563f3e2ba0c8", + "Node name for S&R": "WanVideoEncode" + }, + "widgets_values": [ + false, + 272, + 272, + 144, + 128, + 0, + 0.5 + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 123, + "type": "Note", + "pos": [ + 2024.4696044921875, + -677.442626953125 + ], + "size": [ + 239.09585571289062, + 94.3968734741211 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "The sampling windows are constant size, rounded up and padded with empty if there isn't enough audio" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 37, + "type": "WanVideoEmptyEmbeds", + "pos": [ + 403.56707763671875, + -283.66357421875 + ], + "size": [ + 315, + 126 + ], + "flags": {}, + "order": 38, + "mode": 0, + "inputs": [ + { + "name": "control_embeds", + "shape": 7, + "type": "WANVIDIMAGE_EMBEDS", + "link": null + }, + { + "name": "extra_latents", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "width", + "type": "INT", + "widget": { + "name": "width" + }, + "link": 85 + }, + { + "name": "height", + "type": "INT", + "widget": { + "name": "height" + }, + "link": 86 + }, + { + "name": "num_frames", + "type": "INT", + "widget": { + "name": "num_frames" + }, + "link": 79 + } + ], + "outputs": [ + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "links": [ + 190 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoEmptyEmbeds" + }, + "widgets_values": [ + 832, + 480, + 501 + ] + }, + { + "id": 72, + "type": "WanVideoEncode", + "pos": [ + 432.53973388671875, + -18.44605827331543 + ], + "size": [ + 270, + 242 + ], + "flags": {}, + "order": 43, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "WANVAE", + "link": 198 + }, + { + "name": "image", + "type": "IMAGE", + "link": 202 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "links": [ + 191 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "63d4b6aadaae543f96d101788122563f3e2ba0c8", + "Node name for S&R": "WanVideoEncode" + }, + "widgets_values": [ + false, + 272, + 272, + 144, + 128, + 0, + 1 + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 121, + "type": "GetNode", + "pos": [ + 422.7891845703125, + -93.28303527832031 + ], + "size": [ + 210, + 50 + ], + "flags": { + "collapsed": true + }, + "order": 10, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "WANVAE", + "type": "WANVAE", + "links": [ + 198, + 200 + ] + } + ], + "title": "Get_VAE", + "properties": {}, + "widgets_values": [ + "VAE" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 94, + "type": "VHS_LoadAudio", + "pos": [ + -1128.97119140625, + 360.84234619140625 + ], + "size": [ + 415.6340026855469, + 126 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "audio", + "type": "AUDIO", + "links": [] + }, + { + "name": "duration", + "type": "FLOAT", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "8e4d79471bf1952154768e8435a9300077b534fa", + "Node name for S&R": "VHS_LoadAudio" + }, + "widgets_values": { + "audio_file": "input/weightoftheworld2.mp4", + "seek_seconds": 0, + "duration": 0 + } + }, + { + "id": 66, + "type": "LoadAudio", + "pos": [ + -1127.343505859375, + 149.00762939453125 + ], + "size": [ + 361.2844543457031, + 155.81912231445312 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "AUDIO", + "type": "AUDIO", + "links": [] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "LoadAudio" + }, + "widgets_values": [ + "0321. Alphaville - Big In Japan.mp3", + null, + null + ] + }, + { + "id": 82, + "type": "MelBandRoFormerSampler", + "pos": [ + -442.4711608886719, + 323.1734619140625 + ], + "size": [ + 222.73397827148438, + 46 + ], + "flags": {}, + "order": 39, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "MELROFORMERMODEL", + "link": 106 + }, + { + "name": "audio", + "type": "AUDIO", + "link": 225 + } + ], + "outputs": [ + { + "name": "vocals", + "type": "AUDIO", + "links": [ + 149 + ] + }, + { + "name": "instruments", + "type": "AUDIO", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-MelBandRoFormer", + "ver": "b40e263224778ec417114d91d8b3b39934e30de5", + "Node name for S&R": "MelBandRoFormerSampler" + }, + "widgets_values": [] + }, + { + "id": 98, + "type": "NormalizeAudioLoudness", + "pos": [ + -434.1120300292969, + 442.410400390625 + ], + "size": [ + 270, + 58 + ], + "flags": {}, + "order": 44, + "mode": 0, + "inputs": [ + { + "name": "audio", + "type": "AUDIO", + "link": 149 + } + ], + "outputs": [ + { + "name": "audio", + "type": "AUDIO", + "links": [ + 150 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "90c3bbb6c2e4ff5e05305e765d007d5e58428ce4", + "Node name for S&R": "NormalizeAudioLoudness" + }, + "widgets_values": [ + -23 + ] + }, + { + "id": 64, + "type": "AudioEncoderEncode", + "pos": [ + -73.06085205078125, + 415.04669189453125 + ], + "size": [ + 285.087890625, + 46 + ], + "flags": {}, + "order": 46, + "mode": 0, + "inputs": [ + { + "name": "audio_encoder", + "type": "AUDIO_ENCODER", + "link": 69 + }, + { + "name": "audio", + "type": "AUDIO", + "link": 150 + } + ], + "outputs": [ + { + "name": "AUDIO_ENCODER_OUTPUT", + "type": "AUDIO_ENCODER_OUTPUT", + "links": [ + 189 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "AudioEncoderEncode" + }, + "widgets_values": [] + }, + { + "id": 65, + "type": "AudioEncoderLoader", + "pos": [ + -97.48931884765625, + 258.31158447265625 + ], + "size": [ + 346.756103515625, + 58 + ], + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "AUDIO_ENCODER", + "type": "AUDIO_ENCODER", + "links": [ + 69 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "AudioEncoderLoader" + }, + "widgets_values": [ + "wav2vec_xlsr_53_english_fp32.safetensors" + ] + }, + { + "id": 124, + "type": "MarkdownNote", + "pos": [ + -97.37991333007812, + 112.38502502441406 + ], + "size": [ + 399.6143798828125, + 90.30195617675781 + ], + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "[https://huggingface.co/Comfy-Org/Wan_2.2_ComfyUI_Repackaged/blob/main/split_files/audio_encoders/wav2vec2_large_english_fp16.safetensors](https://huggingface.co/Comfy-Org/Wan_2.2_ComfyUI_Repackaged/blob/main/split_files/audio_encoders/wav2vec2_large_english_fp16.safetensors)" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 81, + "type": "MelBandRoFormerModelLoader", + "pos": [ + -483.87762451171875, + 185.78466796875 + ], + "size": [ + 316.2164001464844, + 58 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "model", + "type": "MELROFORMERMODEL", + "links": [ + 106 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-MelBandRoFormer", + "ver": "b40e263224778ec417114d91d8b3b39934e30de5", + "Node name for S&R": "MelBandRoFormerModelLoader" + }, + "widgets_values": [ + "MelBandRoFormer\\MelBandRoformer_fp16.safetensors" + ] + }, + { + "id": 74, + "type": "ImageResizeKJv2", + "pos": [ + -209.49288940429688, + -372.07373046875 + ], + "size": [ + 270, + 336 + ], + "flags": {}, + "order": 32, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 83 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + }, + { + "name": "width", + "type": "INT", + "widget": { + "name": "width" + }, + "link": 219 + }, + { + "name": "height", + "type": "INT", + "widget": { + "name": "height" + }, + "link": 220 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 201 + ] + }, + { + "name": "width", + "type": "INT", + "links": [ + 85 + ] + }, + { + "name": "height", + "type": "INT", + "links": [ + 86 + ] + }, + { + "name": "mask", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "ImageResizeKJv2" + }, + "widgets_values": [ + 640, + 640, + "lanczos", + "crop", + "0, 0, 0", + "center", + 16, + "cpu", + "Output: 1 x 960 x 640 | 7.03MB" + ] + }, + { + "id": 125, + "type": "SetNode", + "pos": [ + 128.457763671875, + -356.9394836425781 + ], + "size": [ + 210, + 60 + ], + "flags": { + "collapsed": true + }, + "order": 37, + "mode": 0, + "inputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "link": 201 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 202 + ] + } + ], + "title": "Set_reference_image", + "properties": { + "previousName": "reference_image" + }, + "widgets_values": [ + "reference_image" + ], + "color": "#2a363b", + "bgcolor": "#3f5159" + }, + { + "id": 112, + "type": "ImageConcatMulti", + "pos": [ + 2722.126220703125, + 26.577224731445312 + ], + "size": [ + 270, + 150 + ], + "flags": {}, + "order": 57, + "mode": 0, + "inputs": [ + { + "name": "image_1", + "type": "IMAGE", + "link": 243 + }, + { + "name": "image_2", + "shape": 7, + "type": "IMAGE", + "link": 244 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 213 + ] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0" + }, + "widgets_values": [ + 2, + "right", + false, + null + ] + }, + { + "id": 97, + "type": "VHS_VideoCombine", + "pos": [ + 3336.458740234375, + -465.7515869140625 + ], + "size": [ + 940.8292846679688, + 334 + ], + "flags": {}, + "order": 59, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 208 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": 212 + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "330bce6c3c0d47ebdedcc0348d9ab355707b7523", + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "WanVideo2_2_S2V", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "trim_to_audio": false, + "pingpong": false, + "save_output": false, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "WanVideo2_2_S2V_00014-audio.mp4", + "subfolder": "", + "type": "temp", + "format": "video/h264-mp4", + "frame_rate": 16, + "workflow": "WanVideo2_2_S2V_00014.png", + "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo2_2_S2V_00014-audio.mp4" + } + } + } + }, + { + "id": 111, + "type": "ImageResizeKJv2", + "pos": [ + -158.63368225097656, + 1323.3763427734375 + ], + "size": [ + 270, + 336.00006103515625 + ], + "flags": {}, + "order": 47, + "mode": 4, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 173 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 174, + 214 + ] + }, + { + "name": "width", + "type": "INT", + "links": null + }, + { + "name": "height", + "type": "INT", + "links": null + }, + { + "name": "mask", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "ImageResizeKJv2" + }, + "widgets_values": [ + 640, + 640, + "bilinear", + "stretch", + "0, 0, 0", + "center", + 16, + "gpu" + ] + }, + { + "id": 130, + "type": "Reroute", + "pos": [ + 1995.2110595703125, + 1312.697265625 + ], + "size": [ + 75, + 26 + ], + "flags": {}, + "order": 49, + "mode": 0, + "inputs": [ + { + "name": "", + "type": "*", + "link": 214 + } + ], + "outputs": [ + { + "name": "", + "type": "IMAGE", + "links": [ + 243 + ] + } + ], + "properties": { + "showOutputText": false, + "horizontal": false + } + }, + { + "id": 122, + "type": "GetNode", + "pos": [ + 230.59597778320312, + 1699.09423828125 + ], + "size": [ + 210, + 34 + ], + "flags": { + "collapsed": true + }, + "order": 16, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "WANVAE", + "type": "WANVAE", + "links": [ + 199 + ] + } + ], + "title": "Get_VAE", + "properties": {}, + "widgets_values": [ + "VAE" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 133, + "type": "SetNode", + "pos": [ + -1378.2918701171875, + -320.5899658203125 + ], + "size": [ + 210, + 60 + ], + "flags": { + "collapsed": true + }, + "order": 33, + "mode": 0, + "inputs": [ + { + "name": "INT", + "type": "INT", + "link": 216 + } + ], + "outputs": [ + { + "name": "*", + "type": "*", + "links": null + } + ], + "title": "Set_width", + "properties": { + "previousName": "width" + }, + "widgets_values": [ + "width" + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 132, + "type": "INTConstant", + "pos": [ + -1611.5635986328125, + -232.2957763671875 + ], + "size": [ + 210, + 58 + ], + "flags": {}, + "order": 17, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "value", + "type": "INT", + "links": [ + 217 + ] + } + ], + "title": "Height", + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "INTConstant" + }, + "widgets_values": [ + 640 + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 134, + "type": "SetNode", + "pos": [ + -1378.2921142578125, + -205.04428100585938 + ], + "size": [ + 210, + 60 + ], + "flags": { + "collapsed": true + }, + "order": 31, + "mode": 0, + "inputs": [ + { + "name": "INT", + "type": "INT", + "link": 217 + } + ], + "outputs": [ + { + "name": "*", + "type": "*", + "links": null + } + ], + "title": "Set_height", + "properties": { + "previousName": "height" + }, + "widgets_values": [ + "height" + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 73, + "type": "LoadImage", + "pos": [ + -794.7612915039062, + -369.87554931640625 + ], + "size": [ + 274.080078125, + 314 + ], + "flags": {}, + "order": 18, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 83 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "2b.jpg", + "image" + ] + }, + { + "id": 137, + "type": "GetNode", + "pos": [ + -369.8827819824219, + -271.25408935546875 + ], + "size": [ + 210, + 60 + ], + "flags": { + "collapsed": true + }, + "order": 19, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 219 + ] + } + ], + "title": "Get_width", + "properties": {}, + "widgets_values": [ + "width" + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 138, + "type": "GetNode", + "pos": [ + -372.0632019042969, + -212.39117431640625 + ], + "size": [ + 210, + 58 + ], + "flags": { + "collapsed": true + }, + "order": 20, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 220 + ] + } + ], + "title": "Get_height", + "properties": {}, + "widgets_values": [ + "height" + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 107, + "type": "DWPreprocessor", + "pos": [ + -456.9013977050781, + 1322.96728515625 + ], + "size": [ + 270, + 198 + ], + "flags": {}, + "order": 45, + "mode": 4, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 169 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 173 + ] + }, + { + "name": "POSE_KEYPOINT", + "type": "POSE_KEYPOINT", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui_controlnet_aux", + "ver": "1d7cdce8cb771fbc39a432a6338168c12a338ef4", + "Node name for S&R": "DWPreprocessor" + }, + "widgets_values": [ + "disable", + "disable", + "enable", + 640, + "yolox_l.torchscript.pt", + "dw-ll_ucoco_384_bs5.torchscript.pt" + ] + }, + { + "id": 141, + "type": "GetNode", + "pos": [ + -1261.4315185546875, + 737.1270751953125 + ], + "size": [ + 210, + 34 + ], + "flags": { + "collapsed": true + }, + "order": 21, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 223 + ] + } + ], + "title": "Get_width", + "properties": {}, + "widgets_values": [ + "width" + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 142, + "type": "GetNode", + "pos": [ + -1263.6119384765625, + 795.989990234375 + ], + "size": [ + 210, + 34 + ], + "flags": { + "collapsed": true + }, + "order": 22, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 224 + ] + } + ], + "title": "Get_height", + "properties": {}, + "widgets_values": [ + "height" + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 117, + "type": "WanVideoAddS2VEmbeds", + "pos": [ + 842.3013305664062, + -125.87545776367188 + ], + "size": [ + 327.17578125, + 234 + ], + "flags": {}, + "order": 50, + "mode": 0, + "inputs": [ + { + "name": "embeds", + "type": "WANVIDIMAGE_EMBEDS", + "link": 190 + }, + { + "name": "audio_encoder_output", + "shape": 7, + "type": "AUDIO_ENCODER_OUTPUT", + "link": 189 + }, + { + "name": "ref_latent", + "shape": 7, + "type": "LATENT", + "link": 191 + }, + { + "name": "pose_latent", + "shape": 7, + "type": "LATENT", + "link": 192 + }, + { + "name": "vae", + "shape": 7, + "type": "WANVAE", + "link": 200 + } + ], + "outputs": [ + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "links": [ + 194 + ] + }, + { + "name": "audio_frame_count", + "type": "INT", + "links": [ + 195 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "a5621b87391013155b4f688fbe01dba10a8104aa", + "Node name for S&R": "WanVideoAddS2VEmbeds" + }, + "widgets_values": [ + 80, + 1, + 0, + 1, + true + ] + }, + { + "id": 118, + "type": "PreviewAny", + "pos": [ + 1220.6553955078125, + -59.443397521972656 + ], + "size": [ + 226.4835662841797, + 88 + ], + "flags": {}, + "order": 52, + "mode": 0, + "inputs": [ + { + "name": "source", + "type": "*", + "link": 195 + } + ], + "outputs": [], + "title": "Actual total frame count", + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.52", + "Node name for S&R": "PreviewAny" + }, + "widgets_values": [] + }, + { + "id": 139, + "type": "GetNode", + "pos": [ + -1181.0894775390625, + 1225.9921875 + ], + "size": [ + 210, + 58 + ], + "flags": { + "collapsed": true + }, + "order": 23, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 221, + 229 + ] + } + ], + "title": "Get_width", + "properties": {}, + "widgets_values": [ + "width" + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 140, + "type": "GetNode", + "pos": [ + -1180.52197265625, + 1278.0623779296875 + ], + "size": [ + 210, + 50 + ], + "flags": { + "collapsed": true + }, + "order": 24, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 222, + 230 + ] + } + ], + "title": "Get_height", + "properties": {}, + "widgets_values": [ + "height" + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 116, + "type": "VHS_LoadVideo", + "pos": [ + -1224.3709716796875, + 1321.8804931640625 + ], + "size": [ + 247.455078125, + 452.2566223144531 + ], + "flags": {}, + "order": 35, + "mode": 0, + "inputs": [ + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + }, + { + "name": "custom_width", + "type": "INT", + "widget": { + "name": "custom_width" + }, + "link": 229 + }, + { + "name": "custom_height", + "type": "INT", + "widget": { + "name": "custom_height" + }, + "link": 230 + }, + { + "name": "frame_load_cap", + "type": "INT", + "widget": { + "name": "frame_load_cap" + }, + "link": 228 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 188 + ] + }, + { + "name": "frame_count", + "type": "INT", + "links": null + }, + { + "name": "audio", + "type": "AUDIO", + "links": null + }, + { + "name": "video_info", + "type": "VHS_VIDEOINFO", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "8e4d79471bf1952154768e8435a9300077b534fa", + "Node name for S&R": "VHS_LoadVideo" + }, + "widgets_values": { + "video": "weight-world-bones_00003-audio.mp4", + "force_rate": 16, + "custom_width": 0, + "custom_height": 0, + "frame_load_cap": 501, + "skip_first_frames": 0, + "select_every_nth": 1, + "format": "Wan", + "choose video to upload": "image", + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "weight-world-bones_00003-audio.mp4", + "type": "input", + "format": "video/mp4", + "force_rate": 16, + "custom_width": 0, + "custom_height": 0, + "frame_load_cap": 501, + "skip_first_frames": 0, + "select_every_nth": 1 + } + } + } + }, + { + "id": 106, + "type": "VHS_LoadVideo", + "pos": [ + -1112.86865234375, + 593.5896606445312 + ], + "size": [ + 247.455078125, + 451.9747314453125 + ], + "flags": {}, + "order": 34, + "mode": 0, + "inputs": [ + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + }, + { + "name": "custom_width", + "type": "INT", + "widget": { + "name": "custom_width" + }, + "link": 223 + }, + { + "name": "custom_height", + "type": "INT", + "widget": { + "name": "custom_height" + }, + "link": 224 + }, + { + "name": "frame_load_cap", + "type": "INT", + "widget": { + "name": "frame_load_cap" + }, + "link": 227 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [] + }, + { + "name": "frame_count", + "type": "INT", + "links": null + }, + { + "name": "audio", + "type": "AUDIO", + "links": [ + 225, + 226 + ] + }, + { + "name": "video_info", + "type": "VHS_VIDEOINFO", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "8e4d79471bf1952154768e8435a9300077b534fa", + "Node name for S&R": "VHS_LoadVideo" + }, + "widgets_values": { + "video": "weightoftheworld2.mp4", + "force_rate": 16, + "custom_width": 0, + "custom_height": 0, + "frame_load_cap": 501, + "skip_first_frames": 0, + "select_every_nth": 1, + "format": "Wan", + "choose video to upload": "image", + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "weightoftheworld2.mp4", + "type": "input", + "format": "video/mp4", + "force_rate": 16, + "custom_width": 0, + "custom_height": 0, + "frame_load_cap": 501, + "skip_first_frames": 0, + "select_every_nth": 1 + } + } + } + }, + { + "id": 110, + "type": "ImageResizeKJv2", + "pos": [ + -759.1649169921875, + 1320.775146484375 + ], + "size": [ + 270, + 336.00006103515625 + ], + "flags": {}, + "order": 41, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 188 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + }, + { + "name": "width", + "type": "INT", + "widget": { + "name": "width" + }, + "link": 221 + }, + { + "name": "height", + "type": "INT", + "widget": { + "name": "height" + }, + "link": 222 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 169 + ] + }, + { + "name": "width", + "type": "INT", + "links": null + }, + { + "name": "height", + "type": "INT", + "links": null + }, + { + "name": "mask", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "ImageResizeKJv2" + }, + "widgets_values": [ + 640, + 640, + "bilinear", + "crop", + "0, 0, 0", + "center", + 16, + "cpu" + ] + }, + { + "id": 131, + "type": "INTConstant", + "pos": [ + -1611.564208984375, + -356.56207275390625 + ], + "size": [ + 210, + 58 + ], + "flags": {}, + "order": 25, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "value", + "type": "INT", + "links": [ + 216 + ] + } + ], + "title": "Width", + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "INTConstant" + }, + "widgets_values": [ + 640 + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 67, + "type": "WanVideoTextEncodeCached", + "pos": [ + 1491.1046142578125, + -938.5902709960938 + ], + "size": [ + 459.45745849609375, + 393.8887939453125 + ], + "flags": {}, + "order": 26, + "mode": 0, + "inputs": [ + { + "name": "extender_args", + "shape": 7, + "type": "WANVIDEOPROMPTEXTENDER_ARGS", + "link": null + } + ], + "outputs": [ + { + "name": "text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "links": [ + 71 + ] + }, + { + "name": "negative_text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "links": null + }, + { + "name": "positive_prompt", + "type": "STRING", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "a1ca0985ec120ff97e34676a64de19c99767bbd4", + "Node name for S&R": "WanVideoTextEncodeCached" + }, + "widgets_values": [ + "umt5-xxl-enc-bf16.safetensors", + "bf16", + "3D animated scene of a young woman singing melancholically", + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", + "disabled", + true, + "gpu" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 70, + "type": "GetImageSizeAndCount", + "pos": [ + 2367.16748046875, + -395.861572265625 + ], + "size": [ + 190.86483764648438, + 86 + ], + "flags": {}, + "order": 54, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 77 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 237 + ] + }, + { + "label": "640 width", + "name": "width", + "type": "INT", + "links": null + }, + { + "label": "640 height", + "name": "height", + "type": "INT", + "links": null + }, + { + "label": "557 count", + "name": "count", + "type": "INT", + "links": [] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "GetImageSizeAndCount" + }, + "widgets_values": [] + }, + { + "id": 127, + "type": "LazySwitchKJ", + "pos": [ + 3008.5068359375, + -267.5511474609375 + ], + "size": [ + 270, + 78 + ], + "flags": {}, + "order": 58, + "mode": 0, + "inputs": [ + { + "name": "on_false", + "type": "*", + "link": 240 + }, + { + "name": "on_true", + "type": "*", + "link": 213 + } + ], + "outputs": [ + { + "name": "*", + "type": "*", + "links": [ + 208 + ] + } + ], + "title": "Switch: Add pose view", + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "LazySwitchKJ" + }, + "widgets_values": [ + true + ] + }, + { + "id": 143, + "type": "GetImageRangeFromBatch", + "pos": [ + 2622.79833984375, + -397.2756652832031 + ], + "size": [ + 340.3267517089844, + 102 + ], + "flags": { + "collapsed": false + }, + "order": 55, + "mode": 0, + "inputs": [ + { + "name": "images", + "shape": 7, + "type": "IMAGE", + "link": 237 + }, + { + "name": "masks", + "shape": 7, + "type": "MASK", + "link": null + }, + { + "name": "num_frames", + "type": "INT", + "widget": { + "name": "num_frames" + }, + "link": 242 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 239 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "GetImageRangeFromBatch" + }, + "widgets_values": [ + 0, + 501 + ] + }, + { + "id": 129, + "type": "Reroute", + "pos": [ + 2847.140380859375, + 632.76953125 + ], + "size": [ + 75, + 26 + ], + "flags": {}, + "order": 40, + "mode": 0, + "inputs": [ + { + "name": "", + "type": "*", + "link": 226 + } + ], + "outputs": [ + { + "name": "", + "type": "AUDIO", + "links": [ + 212 + ] + } + ], + "properties": { + "showOutputText": false, + "horizontal": false + } + }, + { + "id": 71, + "type": "PrimitiveNode", + "pos": [ + -1613.3883056640625, + -71.24759674072266 + ], + "size": [ + 210, + 82 + ], + "flags": {}, + "order": 27, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "INT", + "type": "INT", + "widget": { + "name": "num_frames" + }, + "links": [ + 79, + 227, + 228, + 242 + ] + } + ], + "title": "num_frames", + "properties": { + "Run widget replace on values": false + }, + "widgets_values": [ + 501, + "fixed" + ] + }, + { + "id": 27, + "type": "WanVideoSampler", + "pos": [ + 1616.490966796875, + -391.5707092285156 + ], + "size": [ + 315, + 999 + ], + "flags": {}, + "order": 51, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 60 + }, + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "link": 194 + }, + { + "name": "text_embeds", + "shape": 7, + "type": "WANVIDEOTEXTEMBEDS", + "link": 71 + }, + { + "name": "samples", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "feta_args", + "shape": 7, + "type": "FETAARGS", + "link": null + }, + { + "name": "context_options", + "shape": 7, + "type": "WANVIDCONTEXT", + "link": null + }, + { + "name": "cache_args", + "shape": 7, + "type": "CACHEARGS", + "link": null + }, + { + "name": "flowedit_args", + "shape": 7, + "type": "FLOWEDITARGS", + "link": null + }, + { + "name": "slg_args", + "shape": 7, + "type": "SLGARGS", + "link": null + }, + { + "name": "loop_args", + "shape": 7, + "type": "LOOPARGS", + "link": null + }, + { + "name": "experimental_args", + "shape": 7, + "type": "EXPERIMENTALARGS", + "link": null + }, + { + "name": "sigmas", + "shape": 7, + "type": "SIGMAS", + "link": null + }, + { + "name": "unianimate_poses", + "shape": 7, + "type": "UNIANIMATE_POSE", + "link": null + }, + { + "name": "fantasytalking_embeds", + "shape": 7, + "type": "FANTASYTALKING_EMBEDS", + "link": null + }, + { + "name": "uni3c_embeds", + "shape": 7, + "type": "UNI3C_EMBEDS", + "link": null + }, + { + "name": "multitalk_embeds", + "shape": 7, + "type": "MULTITALK_EMBEDS", + "link": null + }, + { + "name": "freeinit_args", + "shape": 7, + "type": "FREEINITARGS", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "slot_index": 0, + "links": [ + 105 + ] + }, + { + "name": "denoised_samples", + "type": "LATENT", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "5406a72f62adf4a31a8a0a0e4923cc5288399652", + "Node name for S&R": "WanVideoSampler" + }, + "widgets_values": [ + 4, + 1, + 4, + 45, + "fixed", + true, + "dpm++_sde", + 0, + 1, + false, + "comfy", + 0, + -1, + false + ] + }, + { + "id": 105, + "type": "ColorMatch", + "pos": [ + 2680.017333984375, + -202.95472717285156 + ], + "size": [ + 270, + 126 + ], + "flags": {}, + "order": 56, + "mode": 0, + "inputs": [ + { + "name": "image_ref", + "type": "IMAGE", + "link": 203 + }, + { + "name": "image_target", + "type": "IMAGE", + "link": 239 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 240, + 244 + ] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "ba9153cb06fc77bfd86c36835f1817482e8328a0", + "Node name for S&R": "ColorMatch" + }, + "widgets_values": [ + "mkl", + 1, + true + ] + }, + { + "id": 126, + "type": "GetNode", + "pos": [ + 2435.40625, + -167.38819885253906 + ], + "size": [ + 210.99107360839844, + 60 + ], + "flags": { + "collapsed": true + }, + "order": 28, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 203 + ] + } + ], + "title": "Get_reference_image", + "properties": {}, + "widgets_values": [ + "reference_image" + ], + "color": "#2a363b", + "bgcolor": "#3f5159" + } + ], + "links": [ + [ + 58, + 39, + 0, + 56, + 1, + "BLOCKSWAPARGS" + ], + [ + 60, + 56, + 0, + 27, + 0, + "WANVIDEOMODEL" + ], + [ + 61, + 22, + 0, + 58, + 0, + "WANVIDEOMODEL" + ], + [ + 62, + 58, + 0, + 56, + 0, + "WANVIDEOMODEL" + ], + [ + 64, + 60, + 0, + 58, + 1, + "WANVIDLORA" + ], + [ + 69, + 65, + 0, + 64, + 0, + "AUDIO_ENCODER" + ], + [ + 71, + 67, + 0, + 27, + 2, + "WANVIDEOTEXTEMBEDS" + ], + [ + 77, + 28, + 0, + 70, + 0, + "IMAGE" + ], + [ + 79, + 71, + 0, + 37, + 4, + "INT" + ], + [ + 83, + 73, + 0, + 74, + 0, + "IMAGE" + ], + [ + 85, + 74, + 1, + 37, + 2, + "INT" + ], + [ + 86, + 74, + 2, + 37, + 3, + "INT" + ], + [ + 105, + 27, + 0, + 28, + 1, + "LATENT" + ], + [ + 106, + 81, + 0, + 82, + 0, + "MELROFORMERMODEL" + ], + [ + 129, + 35, + 0, + 22, + 0, + "WANCOMPILEARGS" + ], + [ + 149, + 82, + 0, + 98, + 0, + "AUDIO" + ], + [ + 150, + 98, + 0, + 64, + 1, + "AUDIO" + ], + [ + 169, + 110, + 0, + 107, + 0, + "IMAGE" + ], + [ + 173, + 107, + 0, + 111, + 0, + "IMAGE" + ], + [ + 174, + 111, + 0, + 109, + 1, + "IMAGE" + ], + [ + 188, + 116, + 0, + 110, + 0, + "IMAGE" + ], + [ + 189, + 64, + 0, + 117, + 1, + "AUDIO_ENCODER_OUTPUT" + ], + [ + 190, + 37, + 0, + 117, + 0, + "WANVIDIMAGE_EMBEDS" + ], + [ + 191, + 72, + 0, + 117, + 2, + "LATENT" + ], + [ + 192, + 109, + 0, + 117, + 3, + "LATENT" + ], + [ + 194, + 117, + 0, + 27, + 1, + "WANVIDIMAGE_EMBEDS" + ], + [ + 195, + 117, + 1, + 118, + 0, + "*" + ], + [ + 196, + 38, + 0, + 119, + 0, + "*" + ], + [ + 197, + 120, + 0, + 28, + 0, + "WANVAE" + ], + [ + 198, + 121, + 0, + 72, + 0, + "WANVAE" + ], + [ + 199, + 122, + 0, + 109, + 0, + "WANVAE" + ], + [ + 200, + 121, + 0, + 117, + 4, + "WANVAE" + ], + [ + 201, + 74, + 0, + 125, + 0, + "*" + ], + [ + 202, + 125, + 0, + 72, + 1, + "IMAGE" + ], + [ + 203, + 126, + 0, + 105, + 0, + "IMAGE" + ], + [ + 208, + 127, + 0, + 97, + 0, + "IMAGE" + ], + [ + 212, + 129, + 0, + 97, + 1, + "AUDIO" + ], + [ + 213, + 112, + 0, + 127, + 1, + "*" + ], + [ + 214, + 111, + 0, + 130, + 0, + "*" + ], + [ + 216, + 131, + 0, + 133, + 0, + "*" + ], + [ + 217, + 132, + 0, + 134, + 0, + "*" + ], + [ + 219, + 137, + 0, + 74, + 2, + "INT" + ], + [ + 220, + 138, + 0, + 74, + 3, + "INT" + ], + [ + 221, + 139, + 0, + 110, + 2, + "INT" + ], + [ + 222, + 140, + 0, + 110, + 3, + "INT" + ], + [ + 223, + 141, + 0, + 106, + 2, + "INT" + ], + [ + 224, + 142, + 0, + 106, + 3, + "INT" + ], + [ + 225, + 106, + 2, + 82, + 1, + "AUDIO" + ], + [ + 226, + 106, + 2, + 129, + 0, + "*" + ], + [ + 227, + 71, + 0, + 106, + 4, + "INT" + ], + [ + 228, + 71, + 0, + 116, + 4, + "INT" + ], + [ + 229, + 139, + 0, + 116, + 2, + "INT" + ], + [ + 230, + 140, + 0, + 116, + 3, + "INT" + ], + [ + 237, + 70, + 0, + 143, + 0, + "IMAGE" + ], + [ + 239, + 143, + 0, + 105, + 1, + "IMAGE" + ], + [ + 240, + 105, + 0, + 127, + 0, + "*" + ], + [ + 242, + 71, + 0, + 143, + 2, + "INT" + ], + [ + 243, + 130, + 0, + 112, + 0, + "IMAGE" + ], + [ + 244, + 105, + 0, + 112, + 1, + "IMAGE" + ] + ], + "groups": [ + { + "id": 1, + "title": "Models", + "bounding": [ + -1188.1114501953125, + -1609.4483642578125, + 1896.0758056640625, + 1103.56005859375 + ], + "color": "#88A", + "font_size": 24, + "flags": {} + }, + { + "id": 2, + "title": "Pose input (optional)", + "bounding": [ + -1292.3486328125, + 1142.282958984375, + 1939.6778564453125, + 689.3397827148438 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 3, + "title": "Audio input", + "bounding": [ + -1295.53857421875, + 13.78874397277832, + 1630.400146484375, + 1090.776611328125 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + }, + { + "id": 4, + "title": "Separate vocals from music", + "bounding": [ + -617.63525390625, + 72.55604553222656, + 478.7571105957031, + 458.2166748046875 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.5559917313492586, + "offset": [ + 513.4927365364142, + 697.7427692064874 + ] + }, + "frontendVersion": "1.26.6", + "node_versions": { + "ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd", + "comfy-core": "0.3.26", + "ComfyUI-VideoHelperSuite": "0a75c7958fe320efcb052f1d9f8451fd20c730a8" + }, + "VHS_latentpreview": true, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true + }, + "version": 0.4 +} \ No newline at end of file diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index ce9cc2f..22f39df 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -40,11 +40,11 @@ class FramePackMotioner(nn.Module): 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 - dtype=None, device=None): + ): super().__init__() - self.proj = nn.Conv3d(16, inner_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2), dtype=dtype, device=device) - self.proj_2x = nn.Conv3d(16, inner_dim, kernel_size=(2, 4, 4), stride=(2, 4, 4), dtype=dtype, device=device) - self.proj_4x = nn.Conv3d(16, inner_dim, kernel_size=(4, 8, 8), stride=(4, 8, 8), dtype=dtype, device=device) + 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 @@ -79,9 +79,9 @@ class FramePackMotioner(nn.Module): motion_lat = torch.cat([clean_latents_post, clean_latents_2x, clean_latents_4x], dim=1) - rope_post = rope_embedder.rope_encode(1, lat_height, lat_width, t_start=-1, device=motion_latents.device, dtype=motion_latents.dtype) - rope_2x = rope_embedder.rope_encode(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(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_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 @@ -1714,15 +1714,13 @@ class WanModel(torch.nn.Module): # WanLayerNorm(motioner_dim), # zero_module(nn.Linear(motioner_dim, self.dim))) - self.enable_framepack = enable_framepack + enable_framepack = True if enable_framepack: self.frame_packer = FramePackMotioner( inner_dim=self.dim, num_heads=self.num_heads, zip_frame_buckets=[1, 2, 16], - drop_mode='padd', - device=self.main_device, - dtype=self.dtype) + drop_mode='padd') @staticmethod def _prepare_blockwise_causal_attn_mask( @@ -2139,7 +2137,8 @@ class WanModel(torch.nn.Module): s2v_ref_latent=None, s2v_audio_scale=1.0, s2v_ref_motion=None, - s2v_pose=None + s2v_pose=None, + s2v_motion_frames=[1, 0], ): r""" @@ -2189,17 +2188,15 @@ class WanModel(torch.nn.Module): 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 - #motion_frames=[17, 5] - motion_frames=[1, 0] - s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, motion_frames[0]), s2v_audio_input], dim=-1) + 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[:, motion_frames[1]:].clone() + self.audio_emb_global = audio_emb_global[:, s2v_motion_frames[1]:].clone() else: audio_emb = audio_emb_res - merged_audio_emb = audio_emb[:, motion_frames[1]:, :] + merged_audio_emb = audio_emb[:, s2v_motion_frames[1]:, :] # params device = self.patch_embedding.weight.device @@ -2245,16 +2242,14 @@ class WanModel(torch.nn.Module): ] if s2v_pose is not None: - print("s2v_pose.shape:", s2v_pose.shape) - print("x[0].shape:", x[0].shape) 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]) + original_grid_sizes = grid_sizes.clone() 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) @@ -2302,7 +2297,7 @@ class WanModel(torch.nn.Module): 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)] - 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 @@ -2344,7 +2339,7 @@ class WanModel(torch.nn.Module): s2v_ref_latent.shape[2], s2v_ref_latent.shape[3], s2v_ref_latent.shape[4], - t_start=30, device=x.device, dtype=x.dtype) + t_start=max(30, F + 9), device=x.device, dtype=x.dtype) freqs = torch.cat([freqs, freqs_ref], dim=1) self.cached_freqs = freqs @@ -2746,10 +2741,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( @@ -2796,9 +2791,11 @@ class WanModel(torch.nn.Module): 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 = 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) diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index c778fe1..e878a65 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -133,9 +133,6 @@ 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"timesteps: {timesteps}") - if hasattr(sample_scheduler, 'timesteps'): sample_scheduler.timesteps = timesteps From a053336c1a5f1ac5e5743a8529704fc4cc8ba04c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 28 Aug 2025 18:52:10 +0300 Subject: [PATCH 16/31] Update nodes.py --- nodes.py | 26 ++++++++++++++------------ 1 file changed, 14 insertions(+), 12 deletions(-) diff --git a/nodes.py b/nodes.py index b7ad45b..7d90502 100644 --- a/nodes.py +++ b/nodes.py @@ -3720,18 +3720,19 @@ class WanVideoSampler: device=device) videos_last_frames = ref_motion_image - 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()) + 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 @@ -3758,6 +3759,7 @@ class WanVideoSampler: else: input_motion_latents = None + s2v_pose_slice = None if s2v_pose is not None: s2v_pose_slice = pose_cond_list[r].to(device) From 469d9168dad5a47c4ffe5dcc935bf4cee3a3ac8e Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 28 Aug 2025 19:31:47 +0300 Subject: [PATCH 17/31] Update wanvideo2_2_S2V_framepack_pose_testing.json --- s2v/wanvideo2_2_S2V_framepack_pose_testing.json | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/s2v/wanvideo2_2_S2V_framepack_pose_testing.json b/s2v/wanvideo2_2_S2V_framepack_pose_testing.json index 4008de1..b37e5cb 100644 --- a/s2v/wanvideo2_2_S2V_framepack_pose_testing.json +++ b/s2v/wanvideo2_2_S2V_framepack_pose_testing.json @@ -2804,7 +2804,7 @@ 45, "fixed", true, - "dpm++_sde", + "lcm", 0, 1, false, From a21e4b3210f2fd137793b5d499974364c98f86be Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 28 Aug 2025 20:54:36 +0300 Subject: [PATCH 18/31] Update model.py --- wanvideo/modules/model.py | 46 +++++---------------------------------- 1 file changed, 6 insertions(+), 40 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 22f39df..4c17e30 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1119,7 +1119,7 @@ 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: norm2_x = self.norm2(x) @@ -1677,50 +1677,16 @@ class WanModel(torch.nn.Module): attention_mode=attention_mode ) self.trainable_cond_mask = nn.Embedding(3, self.dim) - self.adain_mode = adain_mode - self.zero_timestep = zero_timestep - - # init motioner - enable_framepack = False - enable_motioner = False - add_last_motion = False - if enable_motioner and enable_framepack: - raise ValueError( - "enable_motioner and enable_framepack are mutually exclusive, please set one of them to False" - ) - self.enable_motioner = enable_motioner - self.add_last_motion = add_last_motion - # if enable_motioner: - # motioner_dim = 2048 - # self.motioner = MotionerTransformers( - # patch_size=(2, 4, 4), - # dim=motioner_dim, - # ffn_dim=motioner_dim, - # freq_dim=256, - # out_dim=16, - # num_heads=16, - # num_layers=13, - # window_size=(-1, -1), - # qk_norm=True, - # cross_attn_norm=False, - # eps=1e-6, - # motion_token_num=4, - # enable_tsm=False, - # motion_stride=4, - # expand_ratio=2, - # trainable_token_pos_emb=False, - # ) - # self.zip_motion_out = torch.nn.Sequential( - # WanLayerNorm(motioner_dim), - # zero_module(nn.Linear(motioner_dim, self.dim))) - - enable_framepack = True - if enable_framepack: + 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( From f9be754980f5b7c2eae496fa91a15fb1f4d037ff Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 00:17:05 +0300 Subject: [PATCH 19/31] cleanup --- wanvideo/modules/model.py | 147 +------------------------------------- 1 file changed, 1 insertion(+), 146 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 4c17e30..3bde4c8 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1677,7 +1677,7 @@ class WanModel(torch.nn.Module): 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, @@ -1838,151 +1838,6 @@ class WanModel(torch.nn.Module): block.to(self.offload_device, non_blocking=self.use_non_blocking) return hints - - def process_motion(self, motion_latents, drop_motion_frames=False): - if drop_motion_frames or motion_latents[0].shape[1] == 0: - return [], [] - self.lat_motion_frames = motion_latents[0].shape[1] - mot = [self.patch_embedding(m.unsqueeze(0)) for m in motion_latents] - batch_size = len(mot) - - mot_remb = [] - flattern_mot = [] - for bs in range(batch_size): - height, width = mot[bs].shape[3], mot[bs].shape[4] - flat_mot = mot[bs].flatten(2).transpose(1, 2).contiguous() - motion_grid_sizes = [[ - torch.tensor([-self.lat_motion_frames, 0, - 0]).unsqueeze(0).repeat(1, 1), - torch.tensor([0, height, width]).unsqueeze(0).repeat(1, 1), - torch.tensor([self.lat_motion_frames, height, - width]).unsqueeze(0).repeat(1, 1) - ]] - motion_rope_emb = rope_precompute( - flat_mot.detach().view(1, flat_mot.shape[1], self.num_heads, - self.dim // self.num_heads), - motion_grid_sizes, - self.freqs, - start=None) - mot_remb.append(motion_rope_emb) - flattern_mot.append(flat_mot) - return flattern_mot, mot_remb - - def process_motion_frame_pack(self, - motion_latents, - drop_motion_frames=False, - add_last_motion=2): - flattern_mot, mot_remb = self.frame_packer(motion_latents, - add_last_motion) - if drop_motion_frames: - return [m[:, :0] for m in flattern_mot - ], [m[:, :0] for m in mot_remb] - else: - return flattern_mot, mot_remb - - def process_motion_transformer_motioner(self, - motion_latents, - drop_motion_frames=False, - add_last_motion=True): - batch_size, height, width = len( - motion_latents), motion_latents[0].shape[2] // self.patch_size[ - 1], motion_latents[0].shape[3] // self.patch_size[2] - - freqs = self.freqs - device = self.patch_embedding.weight.device - if freqs.device != device: - freqs = freqs.to(device) - if self.trainable_token_pos_emb: - token_freqs = self.token_freqs.to(torch.float64) - token_freqs = token_freqs / token_freqs.norm( - dim=-1, keepdim=True) - freqs = [freqs, torch.view_as_complex(token_freqs)] - - if not drop_motion_frames and add_last_motion: - last_motion_latent = [u[:, -1:] for u in motion_latents] - last_mot = [ - self.patch_embedding(m.unsqueeze(0)) for m in last_motion_latent - ] - last_mot = [m.flatten(2).transpose(1, 2) for m in last_mot] - last_mot = torch.cat(last_mot) - gride_sizes = [[ - torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([0, height, - width]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([1, height, - width]).unsqueeze(0).repeat(batch_size, 1) - ]] - else: - last_mot = torch.zeros([batch_size, 0, self.dim], - device=motion_latents[0].device, - dtype=motion_latents[0].dtype) - gride_sizes = [] - - zip_motion = self.motioner(motion_latents) - zip_motion = self.zip_motion_out(zip_motion) - if drop_motion_frames: - zip_motion = zip_motion * 0.0 - zip_motion_grid_sizes = [[ - torch.tensor([-1, 0, 0]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor([ - 0, self.motioner.motion_side_len, self.motioner.motion_side_len - ]).unsqueeze(0).repeat(batch_size, 1), - torch.tensor( - [1 if not self.trainable_token_pos_emb else -1, height, - width]).unsqueeze(0).repeat(batch_size, 1), - ]] - - mot = torch.cat([last_mot, zip_motion], dim=1) - gride_sizes = gride_sizes + zip_motion_grid_sizes - - motion_rope_emb = rope_precompute( - mot.detach().view(batch_size, mot.shape[1], self.num_heads, - self.dim // self.num_heads), - gride_sizes, - freqs, - start=None) - return [m.unsqueeze(0) for m in mot - ], [r.unsqueeze(0) for r in motion_rope_emb] - - def inject_motion(self, - x, - seq_lens, - rope_embs, - mask_input, - motion_latents, - drop_motion_frames=False, - add_last_motion=True): - # inject the motion frames token to the hidden states - if self.enable_motioner: - mot, mot_remb = self.process_motion_transformer_motioner( - motion_latents, - drop_motion_frames=drop_motion_frames, - add_last_motion=add_last_motion) - elif self.enable_framepack: - mot, mot_remb = self.process_motion_frame_pack( - motion_latents, - drop_motion_frames=drop_motion_frames, - add_last_motion=add_last_motion) - else: - mot, mot_remb = self.process_motion( - motion_latents, drop_motion_frames=drop_motion_frames) - - if len(mot) > 0: - x = [torch.cat([u, m], dim=1) for u, m in zip(x, mot)] - seq_lens = seq_lens + torch.tensor([r.size(1) for r in mot], - dtype=torch.long) - rope_embs = [ - torch.cat([u, m], dim=1) for u, m in zip(rope_embs, mot_remb) - ] - mask_input = [ - torch.cat([ - m, 2 * torch.ones([1, u.shape[1] - m.shape[1]], - device=m.device, - dtype=m.dtype) - ], - dim=1) for m, u in zip(mask_input, x) - ] - return x, seq_lens, rope_embs, mask_input def audio_injector_forward(self, block_idx, x, audio_emb, scale=1.0): if block_idx in self.audio_injector.injected_block_id.keys(): From 6978686272ffc2258b0cf3083b394414520163b9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 15:16:02 +0300 Subject: [PATCH 20/31] Update model.py --- wanvideo/modules/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 3bde4c8..756f1c7 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2516,7 +2516,7 @@ class WanModel(torch.nn.Module): self.controlnet.to(self.offload_device) # Asynchronous block offloading with CUDA streams and events - cuda_stream = mm.get_offload_stream(device) + cuda_stream = torch.cuda.Stream(device=device, priority=0) events = [torch.cuda.Event() for _ in self.blocks] swap_start_idx = len(self.blocks) - self.blocks_to_swap if self.blocks_to_swap > 0 else len(self.blocks) From 0c8a883a36dafdf1eb45e9cced3e096978a9fc6f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 15:25:31 +0300 Subject: [PATCH 21/31] Update nodes.py --- nodes.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/nodes.py b/nodes.py index 7d90502..f51fe38 100644 --- a/nodes.py +++ b/nodes.py @@ -36,6 +36,18 @@ offload_device = mm.unet_offload_device() VAE_STRIDE = (4, 8, 8) PATCH_SIZE = (1, 2, 2) +try: + from .gguf.gguf import GGUFParameter +except: + pass + +class MetaParameter(torch.nn.Parameter): + def __new__(cls, dtype, quant_type=None): + data = torch.empty(0, dtype=dtype) + self = torch.nn.Parameter(data, requires_grad=False) + self.quant_type = quant_type + return self + def offload_transformer(transformer): for block in transformer.blocks: block.kv_cache = None @@ -55,6 +67,9 @@ def offload_transformer(transformer): if param.data.is_floating_point(): meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False) setattr(module, attr_name, meta_param) + elif isinstance(param.data, GGUFParameter): + quant_type = getattr(param, 'quant_type', None) + setattr(module, attr_name, MetaParameter(param.data.dtype, quant_type)) else: pass else: From 6a2053c9d1876f24360a15246f5e7f889dcb6ad1 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 15:26:47 +0300 Subject: [PATCH 22/31] Update nodes_model_loading.py --- nodes_model_loading.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 4612e95..61ed516 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -778,7 +778,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, weights = torch.from_numpy(tensor.data.copy()).to(load_device) sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights sd.update(unianimate_sd) - del unianimate_sd + del all_tensors, unianimate_sd if not getattr(transformer, "gguf_patched", False): transformer = _replace_with_gguf_linear( @@ -906,6 +906,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False): # Filter out any LoRA keys containing 'img' if the base model state_dict has no 'img' keys #if not any('img' in k for k in sd.keys()): # lora_sd = {k: v for k, v in lora_sd.items() if 'img' not in k} + if "diffusion_model.patch_embedding.lora_A.weight" in lora_sd: control_lora = True #stand-in LoRA patch From 5f7a5d533b8dc8bb2845e3ab26552b1d440579ae Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 16:59:23 +0300 Subject: [PATCH 23/31] Reduce memory use in the Framepack loop --- nodes.py | 17 ++++++++++++----- 1 file changed, 12 insertions(+), 5 deletions(-) diff --git a/nodes.py b/nodes.py index f51fe38..f48fe67 100644 --- a/nodes.py +++ b/nodes.py @@ -2879,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] @@ -3750,11 +3751,15 @@ class WanVideoSampler: 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") + + mm.soft_empty_cache() + gc.collect() # sample for r in range(s2v_num_repeat): 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 @@ -3801,11 +3806,13 @@ class WanVideoSampler: 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:] @@ -3820,7 +3827,9 @@ class WanVideoSampler: ref_motion_image = videos_last_frames - vae.to(offload_device) + 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: @@ -3837,16 +3846,14 @@ class WanVideoSampler: else: noise_pred, self.cache_state = predict_with_cfg( latent_model_input, - cfg[idx], - text_embeds["prompt_embeds"], + cfg, 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) if bidirectional_sampling: noise_pred_flipped, self.cache_state = predict_with_cfg( latent_model_input_flipped, - cfg[idx], - text_embeds["prompt_embeds"], + cfg, 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) From 876b5bda5ceb510a2c4da0bfe5eb726178e4dc18 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 17:17:37 +0300 Subject: [PATCH 24/31] Update nodes.py --- nodes.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/nodes.py b/nodes.py index f48fe67..87ba25c 100644 --- a/nodes.py +++ b/nodes.py @@ -3751,11 +3751,11 @@ class WanVideoSampler: 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") - - mm.soft_empty_cache() - gc.collect() # 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] From 8f6507fa64121ea3be91e7c15197d275f4bc5f29 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 18:10:18 +0300 Subject: [PATCH 25/31] Update nodes.py --- nodes.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/nodes.py b/nodes.py index 87ba25c..b763078 100644 --- a/nodes.py +++ b/nodes.py @@ -3846,14 +3846,14 @@ class WanVideoSampler: else: noise_pred, self.cache_state = predict_with_cfg( latent_model_input, - cfg, 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, s2v_audio_input=s2v_audio_input) if bidirectional_sampling: noise_pred_flipped, self.cache_state = predict_with_cfg( latent_model_input_flipped, - cfg, 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) From 3a290cfd15b07bf324c9745db81a2023b0f630ae Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 20:41:40 +0300 Subject: [PATCH 26/31] Update nodes_model_loading.py --- nodes_model_loading.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 61ed516..52b12f0 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -822,7 +822,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, if "patch_embedding" in name: dtype_to_use = torch.float32 - load_device = device + load_device = transformer_load_device if block_swap_args is not None: if block_idx is not None: if block_idx >= len(transformer.blocks) - block_swap_args.get("blocks_to_swap", 0): @@ -1343,6 +1343,7 @@ class WanVideoModelLoader: scale_weights.clear() patcher.patches.clear() transformer.patched_linear = False + sd = None else: from .custom_linear import _replace_linear transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights) From 7dd0e1ac617534d70cd1c11483e62716143856b1 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 21:21:24 +0300 Subject: [PATCH 27/31] Update nodes_model_loading.py --- nodes_model_loading.py | 1 + 1 file changed, 1 insertion(+) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 52b12f0..a0c747b 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -824,6 +824,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, load_device = transformer_load_device if block_swap_args is not None: + load_device = device if block_idx is not None: if block_idx >= len(transformer.blocks) - block_swap_args.get("blocks_to_swap", 0): load_device = offload_device From a7cf6d4ed168cf080dda26f430938e7ea97ba628 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 29 Aug 2025 21:33:12 +0300 Subject: [PATCH 28/31] Update nodes.py --- skyreels/nodes.py | 106 +++++++++++++++++++++++++++++----------------- 1 file changed, 66 insertions(+), 40 deletions(-) diff --git a/skyreels/nodes.py b/skyreels/nodes.py index 31bfa3c..9c830ca 100644 --- a/skyreels/nodes.py +++ b/skyreels/nodes.py @@ -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) From 40308f1c7029905c9a1f6a886ce13d23e804e97d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 30 Aug 2025 22:33:32 +0300 Subject: [PATCH 29/31] Allow basic FantasyPortrait + S2V --- s2v/nodes.py | 1 + wanvideo/modules/model.py | 18 +++++++++++------- 2 files changed, 12 insertions(+), 7 deletions(-) diff --git a/s2v/nodes.py b/s2v/nodes.py index 5c18cc6..b008c82 100644 --- a/s2v/nodes.py +++ b/s2v/nodes.py @@ -71,6 +71,7 @@ class WanVideoAddS2VEmbeds: 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] diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 756f1c7..fcc52b7 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -626,7 +626,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) @@ -664,19 +664,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) @@ -692,7 +693,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] @@ -906,6 +907,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={}, @@ -936,6 +938,7 @@ 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] """ + 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] @@ -1107,7 +1110,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, @@ -2458,6 +2461,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, From 10d3a0fbef2dd2865e6aa0fbf64201bf484ea18b Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 30 Aug 2025 23:34:14 +0300 Subject: [PATCH 30/31] Update model.py --- wanvideo/modules/model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index fcc52b7..a27041b 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2520,7 +2520,7 @@ class WanModel(torch.nn.Module): self.controlnet.to(self.offload_device) # Asynchronous block offloading with CUDA streams and events - cuda_stream = torch.cuda.Stream(device=device, priority=0) + cuda_stream = None #torch.cuda.Stream(device=device, priority=0) events = [torch.cuda.Event() for _ in self.blocks] swap_start_idx = len(self.blocks) - self.blocks_to_swap if self.blocks_to_swap > 0 else len(self.blocks) From 26a11c3044d7f2ef999e0b60ecaed08375f0b847 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 31 Aug 2025 18:01:53 +0300 Subject: [PATCH 31/31] Update nodes.py --- multitalk/nodes.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/multitalk/nodes.py b/multitalk/nodes.py index 9831705..2f30f53 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -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")