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] 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