From 79583d4be36a37975644ad7f10d1da9786acd8e2 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 19 Sep 2025 10:14:08 +0300 Subject: [PATCH] Support WanAnimate --- custom_linear.py | 2 +- nodes.py | 634 +++++++++++++++--- nodes_model_loading.py | 7 +- nodes_utility.py | 70 ++ utils.py | 31 +- wanvideo/modules/model.py | 82 ++- wanvideo/modules/wananimate/config.py | 0 wanvideo/modules/wananimate/face_blocks.py | 144 ++++ wanvideo/modules/wananimate/motion_encoder.py | 176 +++++ 9 files changed, 1064 insertions(+), 82 deletions(-) create mode 100644 wanvideo/modules/wananimate/config.py create mode 100644 wanvideo/modules/wananimate/face_blocks.py create mode 100644 wanvideo/modules/wananimate/motion_encoder.py diff --git a/custom_linear.py b/custom_linear.py index 5cabaf0..784163f 100644 --- a/custom_linear.py +++ b/custom_linear.py @@ -13,7 +13,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s module_prefix = prefix + name + "." _replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights) - if isinstance(module, nn.Linear) and "loras" not in module_prefix: + if isinstance(module, nn.Linear) and "loras" not in module_prefix and "face" not in module_prefix: in_features = state_dict[module_prefix + "weight"].shape[1] out_features = state_dict[module_prefix + "weight"].shape[0] if scale_weights is not None: diff --git a/nodes.py b/nodes.py index c14ca37..e794f05 100644 --- a/nodes.py +++ b/nodes.py @@ -80,6 +80,39 @@ def offload_transformer(transformer): mm.soft_empty_cache() gc.collect() + +def init_blockswap(transformer, block_swap_args, model): + if not transformer.patched_linear: + if block_swap_args is not None: + 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) + + 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() + 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) + + class WanVideoEnhanceAVideo: @classmethod def INPUT_TYPES(s): @@ -1031,6 +1064,190 @@ class WanVideoImageToVideoEncode: return (image_embeds,) +# region WanAnimate +class WanVideoAnimateEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "vae": ("WANVAE",), + "width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), + "height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), + "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), + "force_offload": ("BOOLEAN", {"default": True}), + "frame_window_size": ("INT", {"default": 77, "min": 1, "max": 1000, "step": 1, "tooltip": "Number of frames to use for temporal attention window"}), + "colormatch": ( + [ + 'disabled', + 'mkl', + 'hm', + 'reinhard', + 'mvgd', + 'hm-mvgd-hm', + 'hm-mkl-hm', + ], { + "default": 'disabled', "tooltip": "Color matching method to use between the windows" + },), + "pose_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional multiplier for the pose"}), + "face_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional multiplier for the face"}), + }, + "optional": { + "clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}), + "ref_images": ("IMAGE", {"tooltip": "Image to encode"}), + "pose_images": ("IMAGE", {"tooltip": "end frame"}), + "face_images": ("IMAGE", {"tooltip": "end frame"}), + "bg_images": ("IMAGE", {"tooltip": "background images"}), + "mask": ("MASK", {"tooltip": "mask"}), + "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), + } + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + + def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength, + ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None): + + H = height + W = width + + lat_h = H // vae.upsampling_factor + lat_w = W // vae.upsampling_factor + + num_refs = ref_images.shape[0] if ref_images is not None else 0 + + num_frames = ((num_frames - 1) // 4) * 4 + 1 + target_shape = (16, (num_frames - 1) // 4 + 1 + num_refs, lat_h, lat_w) + latent_window_size = ((frame_window_size - 1) // 4) + 1 + + looping = num_frames > frame_window_size + if not looping: + num_frames = num_frames + num_refs * 4 + + vae.to(device) + # Resize and rearrange the input image dimensions + pose_latents = ref_latents = ref_latent = None + if pose_images is not None: + pose_images = pose_images[..., :3] + if pose_images.shape[1] != H or pose_images.shape[2] != W: + resized_pose_images = common_upscale(pose_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1) + else: + resized_pose_images = pose_images.permute(3, 0, 1, 2) # C, T, H, W + resized_pose_images = resized_pose_images * 2 - 1 + pose_latents = vae.encode([resized_pose_images.to(device, vae.dtype)], device,tiled=tiled_vae) + if not looping and pose_latents.shape[2] < latent_window_size: + log.info(f"WanAnimate: Padding pose latents from {pose_latents.shape} to length {latent_window_size}") + pad_len = latent_window_size - pose_latents.shape[2] + pad = torch.zeros(pose_latents.shape[0], pose_latents.shape[1], pad_len, pose_latents.shape[3], pose_latents.shape[4], device=pose_latents.device, dtype=pose_latents.dtype) + pose_latents = torch.cat([pose_latents, pad], dim=2) + print("pose_latents", pose_latents.shape) + del resized_pose_images + + bg_latents = None + if bg_images is not None: + if bg_images.shape[1] != H or bg_images.shape[2] != W: + resized_bg_images = common_upscale(bg_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1) + else: + resized_bg_images = bg_images.permute(3, 0, 1, 2) # C, T, H, W + resized_bg_images = resized_bg_images[:3] * 2 - 1 + if not looping: + bg_latents = vae.encode([resized_bg_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0] + print("bg_latents", bg_latents.shape) + del resized_bg_images + else: + resized_bg_images = resized_bg_images.to(offload_device, dtype=vae.dtype) + + if ref_images is not None: + if ref_images.shape[1] != H or ref_images.shape[2] != W: + resized_ref_images = common_upscale(ref_images.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1) + else: + resized_ref_images = ref_images.permute(3, 0, 1, 2) # C, T, H, W + resized_ref_images = resized_ref_images[:3] * 2 - 1 + + if looping or bg_images is not None: # looping or when using background, encode refs separately + ref_latent = vae.encode([resized_ref_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0] + msk = torch.zeros(4, 1, lat_h, lat_w, device=device, dtype=vae.dtype) + msk[:, :1] = 1 + ref_latent_masked = torch.cat([msk, ref_latent], dim=0) # 4+C 1 H W + msk = torch.zeros(4, (frame_window_size - 1) // 4 + 1, lat_h, lat_w, device=device, dtype=vae.dtype) + + if bg_images is None: + zero_frames = torch.zeros(3, num_frames - num_refs, H, W, device=device, dtype=vae.dtype) + concatenated = torch.cat([resized_ref_images.to(device, dtype=vae.dtype), zero_frames], dim=1) + del zero_frames + ref_latent = vae.encode([concatenated.to(device, vae.dtype)], device,tiled=tiled_vae)[0] + del concatenated + print("ref_latent", ref_latent.shape) + + if mask is None: + ref_mask = torch.zeros(1, num_frames, lat_h, lat_w, device=device, dtype=vae.dtype) + else: + ref_mask = 1 - mask[:num_frames] + ref_mask = common_upscale(ref_mask.unsqueeze(1), lat_w, lat_h, "nearest", "disabled").squeeze(1) + ref_mask = ref_mask.to(vae.dtype).to(device) + ref_mask = ref_mask.unsqueeze(-1).permute(3, 0, 1, 2) # C, T, H, W + + if bg_images is None: + ref_mask[:, :num_refs] = 1 + ref_mask_mask_repeated = torch.repeat_interleave(ref_mask[:, 0:1], repeats=4, dim=1) # T, C, H, W + ref_mask = torch.cat([ref_mask_mask_repeated, ref_mask[:, 1:]], dim=1) + ref_mask = ref_mask.view(1, ref_mask.shape[1] // 4, 4, lat_h, lat_w) # 1, T, C, H, W + ref_mask = ref_mask.movedim(1, 2)[0]# C, T, H, W + + if not looping: + if bg_images is not None: + bg_latents_masked = torch.cat([ref_mask, bg_latents], dim=0) + ref_latent = torch.cat([ref_latent_masked, bg_latents_masked], dim=1) + else: + ref_latent = torch.cat([ref_mask, ref_latent], dim=0) + else: + ref_latent = ref_latent_masked + + if face_images is not None: + face_images = face_images[..., :3] + if face_images.shape[1] != 512 or face_images.shape[2] != 512: + resized_face_images = common_upscale(face_images.movedim(-1, 1), 512, 512, "lanczos", "center").movedim(0, 1) + else: + resized_face_images = face_images.permute(3, 0, 1, 2) # B, C, T, H, W + resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0) + else: + resized_face_images = torch.zeros(1, 3, num_frames, 512, 512, device=device, dtype=torch.float32) + resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype) + + vae.model.clear_cache() + + seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1]) + + if force_offload: + vae.model.to(offload_device) + mm.soft_empty_cache() + gc.collect() + + image_embeds = { + "clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None, + "negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None, + "max_seq_len": seq_len, + "pose_latents": pose_latents, + "bg_images": resized_bg_images if bg_images is not None and looping else None, + "ref_masks": ref_mask if mask is not None and looping else None, + "ref_latent": ref_latent, + "ref_image": resized_ref_images if ref_images is not None else None, + "face_pixels": resized_face_images, + "num_frames": num_frames, + "target_shape": target_shape, + "frame_window_size": frame_window_size, + "lat_h": lat_h, + "lat_w": lat_w, + "vae": vae, + "colormatch": colormatch, + "looping": looping, + "pose_strength": pose_strength, + "face_strength": face_strength, + } + + return (image_embeds,) + class WanVideoEmptyEmbeds: @classmethod def INPUT_TYPES(s): @@ -1813,6 +2030,9 @@ class WanVideoSampler: gguf_reader = model["gguf_reader"] control_lora = model["control_lora"] + vae = image_embeds.get("vae", None) + tiled_vae = image_embeds.get("tiled_vae", False) + transformer_options = patcher.model_options.get("transformer_options", None) merge_loras = transformer_options["merge_loras"] @@ -1990,6 +2210,7 @@ class WanVideoSampler: drop_last = image_embeds.get("drop_last", False) has_ref = image_embeds.get("has_ref", False) + else: #t2v target_shape = image_embeds.get("target_shape", None) if target_shape is None: @@ -2105,11 +2326,12 @@ class WanVideoSampler: phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0) + num_frames = image_embeds.get("num_frames", 0) #HuMo inputs humo_audio = image_embeds.get("humo_audio_emb", None) humo_audio_neg = image_embeds.get("humo_audio_emb_neg", None) humo_reference_count = image_embeds.get("humo_reference_count", 0) - num_frames = image_embeds.get("num_frames", 0) + if humo_audio is not None: from .HuMo.nodes import get_audio_emb_window if not multitalk_sampling: @@ -2140,6 +2362,16 @@ class WanVideoSampler: if not isinstance(humo_audio_cfg_scale, list): humo_audio_cfg_scale = [humo_audio_cfg_scale] * (steps + 1) + # WanAnim inputs + frame_window_size = image_embeds.get("frame_window_size", 77) + wananimate_loop = image_embeds.get("looping", False) + wananim_pose_latents = image_embeds.get("pose_latents", None) + wananim_pose_strength = image_embeds.get("pose_strength", 1.0) + wananim_face_strength = image_embeds.get("face_strength", 1.0) + wananim_face_pixels = image_embeds.get("face_pixels", None) + if image_cond is None: + image_cond = image_embeds.get("ref_latent", None) + latent_video_length = noise.shape[1] # Initialize FreeInit filter if enabled @@ -2349,7 +2581,7 @@ class WanVideoSampler: # vid2vid noise_mask=original_image=None - if samples is not None and not multitalk_sampling: + if samples is not None and not multitalk_sampling and not wananimate_loop: saved_generator_state = samples.get("generator_state", None) if saved_generator_state is not None: seed_g.set_state(saved_generator_state) @@ -2470,38 +2702,7 @@ class WanVideoSampler: gc.collect() #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) - - 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) + init_blockswap(transformer, block_swap_args, model) # Initialize Cache if enabled previous_cache_states = None @@ -2647,7 +2848,8 @@ class WanVideoSampler: 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, s2v_motion_frames=[1, 0], s2v_pose=None, - humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None): + humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None, wananim_pose_latents=None, + wananim_face_pixels=None): nonlocal transformer z = z.to(dtype) autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear) @@ -2821,7 +3023,6 @@ class WanVideoSampler: humo_audio_input_neg = None else: humo_audio_input = humo_audio_input_neg = None - base_params = { 'x': [z], # latent 'y': [image_cond_input] if image_cond_input is not None else None, # image cond @@ -2865,8 +3066,10 @@ class WanVideoSampler: "s2v_audio_scale": s2v_audio_scale if s2v_audio_input is not None else 1.0, # speech-to-video audio scale "s2v_pose": s2v_pose if s2v_pose is not None else None, # speech-to-video pose control "s2v_motion_frames": s2v_motion_frames, # speech-to-video motion frames, - "humo_audio": humo_audio_input, # humo audio input - "humo_audio_scale": humo_audio_scale if humo_audio is not None else 1.0, # humo audio scale + "humo_audio": humo_audio, # humo audio input + "humo_audio_scale": humo_audio_scale if humo_audio is not None else 1, + "wananim_pose_latents": wananim_pose_latents.to(device) if wananim_pose_latents is not None else None, # WanAnimate pose latents + "wananim_face_pixel_values": wananim_face_pixels.to(device, torch.float32) if wananim_face_pixels is not None else None, # WanAnimate face images } batch_size = 1 @@ -3029,7 +3232,7 @@ class WanVideoSampler: from .latent_preview import prepare_callback #custom for tiny VAE previews callback = prepare_callback(patcher, len(timesteps)) - if not multitalk_sampling and not framepack: + if not multitalk_sampling and not framepack and not wananimate_loop: 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") @@ -3117,7 +3320,7 @@ class WanVideoSampler: try: pbar = ProgressBar(len(timesteps)) #region main loop start - for idx, t in enumerate(tqdm(timesteps, disable=multitalk_sampling)): + for idx, t in enumerate(tqdm(timesteps, disable=multitalk_sampling or wananimate_loop)): if flowedit_args is not None: if idx < skip_steps: continue @@ -3411,7 +3614,6 @@ class WanVideoSampler: partial_s2v_audio_input = None if s2v_audio_input is not None: - indices = (torch.arange(4 + 1) - 2) * 1 audio_start = c[0] * 4 audio_end = c[-1] * 4 + 1 center_indices = torch.arange(audio_start, audio_end, 1) @@ -3425,6 +3627,20 @@ class WanVideoSampler: partial_add_cond = None if add_cond is not None: partial_add_cond = add_cond[:, :, c].to(device, dtype) + + partial_wananim_face_pixels = partial_wananim_pose_latents = None + if wananim_face_pixels is not None: + start = c[0] * 4 + end = c[-1] * 4 + center_indices = torch.arange(start, end, 1) + center_indices = torch.clamp(center_indices, min=0, max=wananim_face_pixels.shape[2] - 1) + partial_wananim_face_pixels = wananim_face_pixels[:, :, center_indices].to(device, dtype) + if wananim_pose_latents is not None: + start = c[0] + end = c[-1] + center_indices = torch.arange(start, end, 1) + center_indices = torch.clamp(center_indices, min=0, max=wananim_pose_latents.shape[2] - 1) + partial_wananim_pose_latents = wananim_pose_latents[:, :, center_indices][:, :, :context_frames-1].to(device, dtype) if len(timestep.shape) != 1: partial_timestep = timestep[:, c] @@ -3440,7 +3656,8 @@ class WanVideoSampler: 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, s2v_motion_frames=[1, 0], s2v_pose=partial_s2v_pose, - humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg,) + humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg, + wananim_face_pixels=partial_wananim_face_pixels, wananim_pose_latents=partial_wananim_pose_latents) if cache_args is not None: self.window_tracker.cache_states[window_id] = new_teacache @@ -3461,7 +3678,7 @@ class WanVideoSampler: offloaded = False tiled_vae = image_embeds.get("tiled_vae", False) frame_num = clip_length = image_embeds.get("frame_window_size", 81) - vae = image_embeds.get("vae", None) + clip_embeds = image_embeds.get("clip_context", None) if clip_embeds is not None: clip_embeds = clip_embeds.to(dtype) @@ -3680,37 +3897,8 @@ class WanVideoSampler: load_weights(patcher.model.diffusion_model, patcher.model["sd"], weight_dtype, base_dtype=dtype, transformer_load_device=device, block_swap_args=block_swap_args) elif 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) - #blockswap init - if not transformer.patched_linear: - if block_swap_args is not None: - 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) - - 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() - 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) + init_blockswap(transformer, block_swap_args, device, dtype) # Use the appropriate prompt for this section if len(text_embeds["prompt_embeds"]) > 1: @@ -4072,6 +4260,296 @@ class WanVideoSampler: except: pass return {"video": gen_video_samples}, + # region wananimate loop + elif wananimate_loop: + # calculate frame counts + total_frames = num_frames + overlap = 0 + refert_num = 1 + + real_clip_len = frame_window_size - overlap + last_clip_num = (total_frames - overlap) % real_clip_len + extra = 0 if last_clip_num == 0 else real_clip_len - last_clip_num + target_len = total_frames + extra + target_latent_len = (target_len - 1) // 4 + 2 + latent_window_size = (frame_window_size - 1) // 4 + 1 + + from .utils import tensor_pingpong_pad + + ref_latent = image_embeds.get("ref_latent", None) + ref_images = image_embeds.get("ref_image", None) + ref_masks = image_embeds.get("ref_masks", None) + bg_images = image_embeds.get("bg_images", None) + + pose_input_latents = current_ref_images = face_images = None + #if wananim_pose_latents is not None: + #pose_input_latents = tensor_pingpong_pad(wananim_pose_latents, target_latent_len) + #log.info(f"WanAnimate: Pose input {wananim_pose_latents.shape} padded to shape {pose_input_latents.shape}") + if wananim_face_pixels is not None: + face_images = tensor_pingpong_pad(wananim_face_pixels, target_len) + log.info(f"WanAnimate: Face input {wananim_face_pixels.shape} padded to shape {face_images.shape}") + if ref_masks is not None: + ref_masks_in = tensor_pingpong_pad(ref_masks, target_latent_len) + log.info(f"WanAnimate: Ref masks {ref_masks.shape} padded to shape {ref_masks.shape}") + if bg_images is not None: + bg_images_in = tensor_pingpong_pad(bg_images, target_len) + log.info(f"WanAnimate: BG images {bg_images.shape} padded to shape {bg_images.shape}") + + # if replace_flag: + # bg_images, mask_images = self.prepare_source_for_replace(src_bg_path, src_mask_path) + # bg_images = inputs_padding(bg_images, target_len) + # mask_images = inputs_padding(mask_images, target_len) + + # init variables + offloaded = False + + colormatch = image_embeds.get("colormatch", "disabled") + output_path = image_embeds.get("output_path", "") + offload = image_embeds.get("force_offload", False) + + lat_h, lat_w = noise.shape[2], noise.shape[3] + start = start_latent = img_counter = step_iteration_count = iteration_count = 0 + end = frame_window_size + end_latent = latent_window_size + + estimated_iterations = target_len // frame_window_size + callback = prepare_callback(patcher, estimated_iterations) + log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps") + + # outer WanAnimate loop + gen_video_list = [] + while True: + if start >= total_frames: + break + + mm.soft_empty_cache() + + mask_reft_len = 0 if start == 0 else refert_num + + self.cache_state = [None, None] + + if ref_latent is not None: + vae.to(device) + #ref_latents = vae.encode([ref_images.to(device, vae.dtype)], device,tiled=tiled_vae)[0] + #msk = torch.zeros(4, 1, lat_h, lat_w, device=device, dtype=dtype) + #msk[:, :1] = 1 + #ref_latents = torch.cat([msk, ref_latents], dim=0) # 4+C 1 H W + if ref_masks is not None: + msk = ref_masks_in[:, start_latent:end_latent].to(device, dtype) + if msk.shape[1] < latent_window_size: + log.info(f"WanAnimate: Padding ref masks from {msk.shape} to length {latent_window_size}") + pad_length = latent_window_size - msk.shape[1] + last_frame = msk[:, -1:].repeat(1, pad_length, 1, 1) + msk = torch.cat([msk, last_frame], dim=1) + else: + msk = torch.zeros(4, latent_window_size, lat_h, lat_w, device=device, dtype=dtype) + if bg_images is not None: + bg_image_slice = bg_images_in[:, start:end].to(device) + else: + bg_image_slice = torch.zeros(3, frame_window_size-mask_reft_len, lat_h * 8, lat_w * 8, device=device, dtype=vae.dtype) + if mask_reft_len == 0: + temporal_ref_latents = vae.encode([bg_image_slice], device,tiled=tiled_vae)[0] + else: + concatenated = torch.cat([current_ref_images.to(device, dtype=vae.dtype), bg_image_slice[:, mask_reft_len:]], dim=1) + temporal_ref_latents = vae.encode([concatenated.to(device, vae.dtype)], device,tiled=tiled_vae)[0] + msk[:, :mask_reft_len] = 1 + + vae.model.clear_cache() + vae.to(offload_device) + + temporal_ref_latents = torch.cat([msk, temporal_ref_latents], dim=0) # 4+C T H W + image_cond_in = torch.cat([ref_latent, temporal_ref_latents], dim=1) # 4+C T+trefs H W + + noise = torch.randn(16, latent_window_size + 1, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device) + seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1]) + + pose_input_slice = None + if wananim_pose_latents is not None: + pose_input_slice = wananim_pose_latents[:, :, start_latent:end_latent].to(device, dtype) + # Pad if slice is too short + if pose_input_slice.shape[2] < latent_window_size: + log.info(f"WanAnimate: Padding pose latents from {pose_input_slice.shape} to length {latent_window_size}") + pad_len = latent_window_size - pose_input_slice.shape[2] + pad = torch.zeros(pose_input_slice.shape[0], pose_input_slice.shape[1], pad_len, pose_input_slice.shape[3], pose_input_slice.shape[4], device=pose_input_slice.device, dtype=pose_input_slice.dtype) + pose_input_slice = torch.cat([pose_input_slice, pad], dim=2) + pose_input_slice = pose_input_slice.to(device, dtype) + + if samples is not None: + input_samples = samples["samples"].squeeze(0).to(noise) + # Check if we have enough frames in input_samples + if latent_end_idx > input_samples.shape[1]: + # We need more frames than available - pad the input_samples at the end + pad_length = latent_end_idx - input_samples.shape[1] + last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1) + input_samples = torch.cat([input_samples, last_frame], dim=1) + input_samples = input_samples[:, latent_start_idx:latent_end_idx] + if noise_mask is not None: + original_image = input_samples.to(device) + + assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}" + + if add_noise_to_samples: + latent_timestep = timesteps[0] + noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples + else: + noise = input_samples + + # diff diff prep + noise_mask = samples.get("noise_mask", None) + if noise_mask is not None: + if len(noise_mask.shape) == 4: + noise_mask = noise_mask.squeeze(1) + if noise_mask.shape[0] < noise.shape[1]: + noise_mask = noise_mask.repeat(noise.shape[1] // noise_mask.shape[0], 1, 1) + else: + noise_mask = noise_mask[latent_start_idx:latent_end_idx] + noise_mask = torch.nn.functional.interpolate( + noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W] + size=(noise.shape[1], noise.shape[2], noise.shape[3]), + mode='trilinear', + align_corners=False + ).repeat(1, noise.shape[0], 1, 1, 1) + + thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps) + thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device) + masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds + + sample_scheduler, timesteps,_,_ = get_scheduler(scheduler, total_steps, start_step, end_step, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas) + + # sample videos + latent = noise + + if offloaded: + # 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) + elif 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) + #blockswap init + init_blockswap(transformer, block_swap_args, model) + + # Use the appropriate prompt for this section + if len(text_embeds["prompt_embeds"]) > 1: + prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1) + positive = [text_embeds["prompt_embeds"][prompt_index]] + log.info(f"Using prompt index: {prompt_index}") + else: + positive = text_embeds["prompt_embeds"] + + # uni3c slices + if uni3c_embeds is not None: + vae.to(device) + # Pad original_images if needed + num_frames = original_images.shape[2] + required_frames = audio_end_idx - audio_start_idx + if audio_end_idx > num_frames: + pad_len = audio_end_idx - num_frames + last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1) + padded_images = torch.cat([original_images, last_frame], dim=2) + else: + padded_images = original_images + render_latent = vae.encode( + padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype), + device=device, tiled=tiled_vae + ).to(dtype) + vae.model.clear_cache() + vae.to(offload_device) + pcd_data['render_latent'] = render_latent + + mm.soft_empty_cache() + gc.collect() + # inner WanAnimate sampling loop + sampling_pbar = tqdm(total=len(timesteps), desc=f"Frames {start}-{end}", position=0, leave=True) + for i in range(len(timesteps)): + timestep = timesteps[i] + latent_model_input = latent.to(device) + + noise_pred, self.cache_state = predict_with_cfg( + latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"], + timestep, i, cache_state=self.cache_state, + image_cond = image_cond_in, + wananim_face_pixels=face_images[:, :, start:end].to(device, torch.float32) if face_images is not None else None, + wananim_pose_latents=pose_input_slice + ) + 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, estimated_iterations*(len(timesteps))) + del callback_latent + + sampling_pbar.update(1) + step_iteration_count += 1 + + latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0) + del noise_pred, latent_model_input, timestep + + # differential diffusion inpaint + if masks is not None: + if i < len(timesteps) - 1: + image_latent = add_noise(original_image.to(device), noise.to(device), timesteps[i+1]) + mask = masks[i].to(latent) + latent = image_latent * mask + latent * (1-mask) + + del noise + if offload: + offload_transformer(transformer) + offloaded = True + + vae.to(device) + videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu() + del latent + vae.model.clear_cache() + vae.to(offload_device) + + sampling_pbar.close() + + # optional color correction + if colormatch != "disabled": + videos = videos.permute(1, 2, 3, 0).float().numpy() + from color_matcher import ColorMatcher + cm = ColorMatcher() + cm_result_list = [] + for img in videos: + cm_result = cm.transfer(src=img, ref=ref_images.permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch) + cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype)) + videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2) + + current_ref_images = videos[:, -1:].clone().detach() + + # optionally save generated samples to disk + if output_path: + video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8') + num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num + log.info(f"Saving {num_frames_to_save} generated frames to {output_path}") + start_idx = 0 if is_first_clip else cur_motion_frames_num + for i in range(start_idx, video_np.shape[0]): + im = Image.fromarray(video_np[i]) + im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png")) + img_counter += 1 + else: + gen_video_list.append(videos) + + del videos + + iteration_count += 1 + start += frame_window_size + end += frame_window_size + start_latent += latent_window_size + end_latent += latent_window_size + + if not output_path: + gen_video_samples = torch.cat(gen_video_list, dim=1) + else: + gen_video_samples = torch.zeros(3, 1, 64, 64) # dummy output + + 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), "output_path": output_path}, #region normal inference else: @@ -4081,7 +4559,9 @@ class WanVideoSampler: 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, multitalk_audio_embeds=multitalk_audio_embeds, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input, - humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg) + humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg, + wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents, + ) if bidirectional_sampling: noise_pred_flipped, self.cache_state = predict_with_cfg( latent_model_input_flipped, @@ -4188,6 +4668,8 @@ class WanVideoSampler: latent = latent[:,:-phantom_latents.shape[1]] if humo_reference_count > 0: latent = latent[:,:-humo_reference_count] + if wananim_pose_latents is not None: + latent = latent[:, 1:] cache_states = None if cache_args is not None: @@ -4469,7 +4951,8 @@ NODE_CLASS_MAPPINGS = { "WanVideoAddControlEmbeds": WanVideoAddControlEmbeds, "WanVideoAddMTVMotion": WanVideoAddMTVMotion, "WanVideoRoPEFunction": WanVideoRoPEFunction, - "WanVideoAddPusaNoise": WanVideoAddPusaNoise + "WanVideoAddPusaNoise": WanVideoAddPusaNoise, + "WanVideoAnimateEmbeds": WanVideoAnimateEmbeds, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -4506,4 +4989,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoAddMTVMotion": "WanVideo MTV Crafter Motion", "WanVideoRoPEFunction": "WanVideo RoPE Function", "WanVideoAddPusaNoise": "WanVideo Add Pusa Noise", + "WanVideoAnimateEmbeds": "WanVideo Animate Embeds", } diff --git a/nodes_model_loading.py b/nodes_model_loading.py index ccae150..f8315e3 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -761,7 +761,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", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob"} + "adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "motion_encoder"} param_count = sum(1 for _ in transformer.named_parameters()) pbar = ProgressBar(param_count) cnt = 0 @@ -851,7 +851,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, dtype_to_use = weight_dtype if sd[name.replace("_orig_mod.", "")].dtype == weight_dtype else dtype_to_use if "modulation" in name or "norm" in name or "bias" in name or "img_emb" in name: dtype_to_use = base_dtype - if "patch_embedding" in name: + if "patch_embedding" in name or "motion_encoder" in name or "face_encoder" in name: dtype_to_use = torch.float32 load_device = transformer_load_device @@ -1116,6 +1116,7 @@ class WanVideoModelLoader: ffn2_dim = sd["blocks.0.ffn.2.weight"].shape[1] is_humo = "audio_proj.audio_proj_glob_1.layer.weight" in sd + is_wananimate = "pose_patch_embedding.weight" in sd model_type = "t2v" if "audio_injector.injector.0.k.weight" in sd: @@ -1219,6 +1220,7 @@ class WanVideoModelLoader: "rope_func": "comfy", "main_device": device, "offload_device": offload_device, + "dtype": base_dtype, "teacache_coefficients": teacache_coefficients_map[model_variant], "magcache_ratios": magcache_ratios_map[model_variant], "vace_layers": vace_layers, @@ -1232,6 +1234,7 @@ class WanVideoModelLoader: "cond_dim": sd["cond_encoder.weight"].shape[1] if "cond_encoder.weight" in sd else 0, "zero_timestep": model_type == "s2v", "humo_audio": is_humo, + "is_wananimate": is_wananimate, } diff --git a/nodes_utility.py b/nodes_utility.py index 81f8fec..76fc512 100644 --- a/nodes_utility.py +++ b/nodes_utility.py @@ -2,6 +2,7 @@ import torch import numpy as np from comfy.utils import common_upscale from .utils import log +from einops import rearrange try: from server import PromptServer @@ -476,6 +477,73 @@ class WanVideoPassImagesFromSamples: video.clamp_(-1.0, 1.0) video.add_(1.0).div_(2.0) return video.cpu().float(), samples.get("output_path", "") + + +class FaceMaskFromPoseKeypoints: + @classmethod + def INPUT_TYPES(s): + input_types = { + "required": { + "pose_kps": ("POSE_KEYPOINT",), + "person_index": ("INT", {"default": 0, "min": 0, "max": 100, "step": 1, "tooltip": "Index of the person to start with"}), + } + } + return input_types + RETURN_TYPES = ("MASK",) + FUNCTION = "createmask" + CATEGORY = "ControlNet Preprocessors/Pose Keypoint Postprocess" + + def createmask(self, pose_kps, person_index): + pose_frames = pose_kps + prev_center = None + np_frames = [] + for i, pose_frame in enumerate(pose_frames): + selected_idx, prev_center = self.select_closest_person(pose_frame, person_index if i == 0 else prev_center) + np_frames.append(self.draw_kps(pose_frame, selected_idx)) + np_frames = np.stack(np_frames, axis=0) + tensor = torch.from_numpy(np_frames).float() / 255. + print("tensor.shape:", tensor.shape) + tensor = tensor[:, :, :, 0] + return (tensor,) + + def select_closest_person(self, pose_frame, prev_center_or_index): + people = pose_frame["people"] + if not people: + return -1, None + centers = [] + for person in people: + kps = np.array(person["face_keypoints_2d"]) + n = len(kps) // 3 + facial_kps = rearrange(kps, "(n c) -> n c", n=n, c=3)[:, :2] + center = facial_kps.mean(axis=0) + centers.append(center) + if isinstance(prev_center_or_index, (int, np.integer)): + # First frame: use person_index + idx = prev_center_or_index if 0 <= prev_center_or_index < len(people) else 0 + return idx, centers[idx] + else: + # Find closest to previous center + prev_center = np.array(prev_center_or_index) + dists = [np.linalg.norm(center - prev_center) for center in centers] + idx = int(np.argmin(dists)) + return idx, centers[idx] + + def draw_kps(self, pose_frame, person_index): + import cv2 + width, height = pose_frame["canvas_width"], pose_frame["canvas_height"] + canvas = np.zeros((height, width, 3), dtype=np.uint8) + people = pose_frame["people"] + if person_index < 0 or person_index >= len(people): + return canvas # Out of bounds, return blank + person = people[person_index] + n = len(person["face_keypoints_2d"]) // 3 + facial_kps = rearrange(np.array(person["face_keypoints_2d"]), "(n c) -> n c", n=n, c=3)[:, :2] + facial_kps = facial_kps.astype(np.int32) + part_color = (255, 255, 255) + + outer_contour = facial_kps[:17] + cv2.fillPoly(canvas, pts=[outer_contour], color=part_color) + return canvas NODE_CLASS_MAPPINGS = { "WanVideoImageResizeToClosest": WanVideoImageResizeToClosest, @@ -488,6 +556,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoSigmaToStep": WanVideoSigmaToStep, "NormalizeAudioLoudness": NormalizeAudioLoudness, "WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples, + "FaceMaskFromPoseKeypoints": FaceMaskFromPoseKeypoints, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest", @@ -500,4 +569,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSigmaToStep": "WanVideo Sigma To Step", "NormalizeAudioLoudness": "Normalize Audio Loudness", "WanVideoPassImagesFromSamples": "WanVideo Pass Images From Samples", + "FaceMaskFromPoseKeypoints": "Face Mask From Pose Keypoints", } \ No newline at end of file diff --git a/utils.py b/utils.py index ac58556..a8c8be0 100644 --- a/utils.py +++ b/utils.py @@ -3,6 +3,7 @@ import torch import logging import math from tqdm import tqdm +from copy import deepcopy import types, collections from comfy.utils import ProgressBar, copy_to_param, set_attr_param from comfy.model_patcher import get_key_weight, string_to_seed @@ -538,4 +539,32 @@ def get_raag_guidance(noise_pred_cond, noise_pred_uncond, w_max, alpha=1.0, eps= ratio = norm_delta / (norm_uncond + eps) ratio_mean = ratio.mean().item() adaptive_w = 1.0 + (w_max - 1.0) * math.exp(-alpha * ratio_mean) - return adaptive_w \ No newline at end of file + return adaptive_w + +def tensor_pingpong_pad(video, target_len): + """ + Pads a video tensor along the frame dimension (dim=2) in a ping-pong fashion. + video: torch.Tensor of shape [B, C, F, H, W] + target_len: desired number of frames + Returns: padded tensor of shape [B, C, target_len, H, W] + """ + in_dims = len(video.shape) + if in_dims == 4: + video = video.unsqueeze(0) + B, C, F, H, W = video.shape + idx = 0 + flip = False + indices = [] + while len(indices) < target_len: + indices.append(idx) + if flip: + idx -= 1 + else: + idx += 1 + if idx == 0 or idx == F - 1: + flip = not flip + indices = indices[:target_len] + padded_video = video[:, :, indices, :, :] + if in_dims == 4: + padded_video = padded_video.squeeze(0) + return padded_video \ No newline at end of file diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index b27e043..d5befae 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -584,7 +584,7 @@ class LoRALinearLayer(nn.Module): down_hidden_states = self.down(hidden_states.to(dtype)) up_hidden_states = self.up(down_hidden_states) * self.strength return up_hidden_states.to(orig_dtype) - + #region crossattn class WanT2VCrossAttention(WanSelfAttention): @@ -1445,6 +1445,7 @@ class WanModel(torch.nn.Module): rope_func='comfy', main_device=torch.device('cuda'), offload_device=torch.device('cpu'), + dtype=torch.float16, teacache_coefficients=[], magcache_ratios=[], vace_layers=None, @@ -1464,6 +1465,9 @@ class WanModel(torch.nn.Module): audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39], zero_timestep=False, humo_audio=False, + # WanAnimate + is_wananimate=False, + motion_encoder_dim=512, ): r""" Initialize the diffusion model backbone. @@ -1575,6 +1579,8 @@ class WanModel(torch.nn.Module): self.humo_audio = humo_audio + self.motion_encoder_dim = motion_encoder_dim + # embeddings self.patch_embedding = nn.Conv3d( in_dim, dim, kernel_size=patch_size, stride=patch_size) @@ -1714,7 +1720,23 @@ class WanModel(torch.nn.Module): from ...HuMo.audio_proj import AudioProjModel self.audio_proj = AudioProjModel(seq_len=8, blocks=5, channels=1280, intermediate_dim=512, output_dim=1536, context_tokens=16) - + # WanAnimate + self.motion_encoder = self.pose_patch_embedding = self.face_encoder = self.face_adapter = None + if is_wananimate: + from .wananimate.motion_encoder import MotionExtractor + from .wananimate.face_blocks import FaceEncoder, FaceAdapter + self.pose_patch_embedding = nn.Conv3d(16, dim, kernel_size=patch_size, stride=patch_size) + self.motion_encoder = MotionExtractor() + self.face_adapter = FaceAdapter( + num_heads=self.num_heads, + feature_dim=self.dim, + num_adapter_layers=self.num_layers // 5, + ) + self.face_encoder = FaceEncoder( + in_dim=motion_encoder_dim, + out_dim=self.dim, + num_heads=4, + ) def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None, prefetch_blocks=0, block_swap_debug=False): # Clamp blocks_to_swap to valid range @@ -1850,6 +1872,44 @@ class WanModel(torch.nn.Module): return x + def wananimate_pose_embedding(self, x, pose_latents, strength=1.0): + pose_latents = [self.pose_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in pose_latents] + for x_, pose_latents_ in zip(x, pose_latents): + x_[:, :, 1:].add_(pose_latents_, alpha=strength) + return x + + + def wananimate_face_embedding(self, face_pixel_values): + b,c,T,h,w = face_pixel_values.shape + face_pixel_values = rearrange(face_pixel_values, "b c t h w -> (b t) c h w") + + encode_bs = 8 + face_pixel_values_tmp = [] + self.motion_encoder.to(self.main_device) + for i in range(math.ceil(face_pixel_values.shape[0]/encode_bs)): + face_pixel_values_tmp.append(self.motion_encoder(face_pixel_values[i*encode_bs:(i+1)*encode_bs])) + del face_pixel_values + self.motion_encoder.to(self.offload_device) + + motion_vec = rearrange(torch.cat(face_pixel_values_tmp), "(b t) c -> b t c", t=T) + del face_pixel_values_tmp + self.face_encoder.to(self.main_device) + motion_vec = self.face_encoder(motion_vec) + self.face_encoder.to(self.offload_device) + + B, L, H, C = motion_vec.shape + pad_face = torch.zeros(B, 1, H, C, device=motion_vec.device, dtype=motion_vec.dtype) + return torch.cat([pad_face, motion_vec], dim=1) + + + def wananimate_forward(self, block_idx, x, motion_vec, strength=1.0, motion_masks=None): + if block_idx % 5 == 0: + adapter_args = [x, motion_vec, motion_masks] + residual_out = self.face_adapter.fuser_blocks[block_idx // 5](*adapter_args) + return x.add(residual_out, alpha=strength) + return x + + def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=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]) @@ -1949,6 +2009,10 @@ class WanModel(torch.nn.Module): s2v_motion_frames=[1, 0], humo_audio=None, humo_audio_scale=1.0, + wananim_pose_latents=None, + wananim_face_pixel_values=None, + wananim_pose_strength=1.0, + wananim_face_strength=1.0, ): r""" @@ -2045,10 +2109,20 @@ class WanModel(torch.nn.Module): self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x ] - + + # WanAnimate + motion_vec = None + if wananim_face_pixel_values is not None: + motion_vec = self.wananimate_face_embedding(wananim_face_pixel_values).to(x[0].dtype) + + if wananim_pose_latents is not None: + x = self.wananimate_pose_embedding(x, wananim_pose_latents, strength=wananim_pose_strength) + + # s2v pose embedding if s2v_pose is not None: x[0] = x[0] + self.cond_encoder(s2v_pose.to(self.cond_encoder.weight.dtype)).to(x[0].dtype) + # Fun camera 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)] @@ -2550,6 +2624,8 @@ class WanModel(torch.nn.Module): x, x_ip = block(x, x_ip=x_ip, **kwargs) #run block if self.audio_injector is not None and s2v_audio_input is not None: x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v + if self.motion_encoder is not None and motion_vec is not None: + x = self.wananimate_forward(b, x, motion_vec, strength=wananim_face_strength) if self.block_swap_debug: compute_end = time.perf_counter() compute_time = compute_end - compute_start diff --git a/wanvideo/modules/wananimate/config.py b/wanvideo/modules/wananimate/config.py new file mode 100644 index 0000000..e69de29 diff --git a/wanvideo/modules/wananimate/face_blocks.py b/wanvideo/modules/wananimate/face_blocks.py new file mode 100644 index 0000000..0627fa8 --- /dev/null +++ b/wanvideo/modules/wananimate/face_blocks.py @@ -0,0 +1,144 @@ +from torch import nn +import torch +from einops import rearrange +import torch.nn.functional as F +from ..attention import attention + +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 FaceEncoder(nn.Module): + def __init__(self, in_dim: int, out_dim: int, num_heads: int, dtype=None, device=None): + super().__init__() + + self.num_heads = num_heads + self.conv1_local = CausalConv1d(in_dim, 1024 * num_heads, 3, stride=1) + self.norm1 = nn.LayerNorm(1024, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype) + self.act = nn.SiLU() + self.conv2 = CausalConv1d(1024, 1024, 3, stride=2) + self.conv3 = CausalConv1d(1024, 1024, 3, stride=2) + + self.out_proj = nn.Linear(1024, out_dim) + + self.norm2 = nn.LayerNorm(1024, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype) + self.norm3 = nn.LayerNorm(1024, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype) + + self.padding_tokens = nn.Parameter(torch.zeros(1, 1, 1, out_dim)) + + def forward(self, x): + x = rearrange(x, "b t c -> b c t") + b = x.shape[0] + + 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 = self.out_proj(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) + + return torch.cat([x, padding], dim=-2) + + +class RMSNorm(nn.Module): + def __init__(self, dim, elementwise_affine=True, eps=1e-6, device=None, dtype=None): + super().__init__() + self.eps = eps + if elementwise_affine: + self.weight = nn.Parameter(torch.ones(dim, device=device, dtype=dtype)) + + def _norm(self, x): + return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) + + def forward(self, x): + output = self._norm(x.float()).type_as(x) + if hasattr(self, "weight"): + output = output * self.weight + return output + + +class FaceAdapter(nn.Module): + def __init__(self, feature_dim, num_heads, num_adapter_layers=1, dtype=None, device=None): + super().__init__() + self.fuser_blocks = nn.ModuleList([FaceBlock(feature_dim, num_heads, device=device, dtype=dtype) for _ in range(num_adapter_layers)]) + + def forward( + self, + x: torch.Tensor, + motion_embed: torch.Tensor, + idx: int, + ) -> torch.Tensor: + + return self.fuser_blocks[idx](x, motion_embed) + + +class FaceBlock(nn.Module): + def __init__(self, feature_dim, num_heads, dtype=None, device=None): + super().__init__() + + self.feature_dim = feature_dim + self.num_heads = num_heads + head_dim = feature_dim // num_heads + + self.linear1_kv = nn.Linear(feature_dim, feature_dim * 2, device=device, dtype=dtype) + self.linear1_q = nn.Linear(feature_dim, feature_dim, device=device, dtype=dtype) + self.linear2 = nn.Linear(feature_dim, feature_dim, device=device, dtype=dtype) + + self.q_norm = (RMSNorm(head_dim, elementwise_affine=True, eps=1e-6, device=device, dtype=dtype)) + self.k_norm = (RMSNorm(head_dim, elementwise_affine=True, eps=1e-6, device=device, dtype=dtype)) + + self.pre_norm_feat = nn.LayerNorm(feature_dim, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype) + self.pre_norm_motion = nn.LayerNorm(feature_dim, elementwise_affine=False, eps=1e-6, device=device, dtype=dtype) + + + def forward(self, x, motion_vec, motion_mask=None): + B, T, N, C = motion_vec.shape + + x_motion = self.pre_norm_motion(motion_vec) + x_feat = self.pre_norm_feat(x) + + kv = self.linear1_kv(x_motion) + q = self.linear1_q(x_feat) + + k, v = rearrange(kv, "B L N (K H D) -> K B L N H D", K=2, H=self.num_heads) + q = rearrange(q, "B S (H D) -> B S H D", H=self.num_heads) + + q = self.q_norm(q).to(v) + k = self.k_norm(k).to(v) + + k = rearrange(k, "B L N H D -> (B L) N H D") + v = rearrange(v, "B L N H D -> (B L) N H D") + q = rearrange(q, "B (L S) H D -> (B L) S H D", L=T) + + attn = attention(q, k, v) + attn = attn.reshape(attn.shape[0], attn.shape[1], -1) + attn = rearrange(attn, "(B L) S C -> B (L S) C", L=T) + output = self.linear2(attn) + + if motion_mask is not None: + output = output * rearrange(motion_mask, "B T H W -> B (T H W)").unsqueeze(-1) + + return output \ No newline at end of file diff --git a/wanvideo/modules/wananimate/motion_encoder.py b/wanvideo/modules/wananimate/motion_encoder.py new file mode 100644 index 0000000..c765597 --- /dev/null +++ b/wanvideo/modules/wananimate/motion_encoder.py @@ -0,0 +1,176 @@ +import torch +from torch.nn import functional as F +import math + +# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/ops/upfirdn2d/upfirdn2d.py#L162 +def upfirdn2d_native(input, kernel, up_x, up_y, down_x, down_y, pad_x0, pad_x1, pad_y0, pad_y1): + _, minor, in_h, in_w = input.shape + kernel_h, kernel_w = kernel.shape + + out = input.view(-1, minor, in_h, 1, in_w, 1) + out = F.pad(out, [0, up_x - 1, 0, 0, 0, up_y - 1, 0, 0]) + out = out.view(-1, minor, in_h * up_y, in_w * up_x) + + out = F.pad(out, [max(pad_x0, 0), max(pad_x1, 0), max(pad_y0, 0), max(pad_y1, 0)]) + out = out[:, :, max(-pad_y0, 0): out.shape[2] - max(-pad_y1, 0), max(-pad_x0, 0): out.shape[3] - max(-pad_x1, 0)] + + out = out.reshape([-1, 1, in_h * up_y + pad_y0 + pad_y1, in_w * up_x + pad_x0 + pad_x1]) + w = torch.flip(kernel, [0, 1]).view(1, 1, kernel_h, kernel_w) + out = F.conv2d(out, w) + out = out.reshape(-1, minor, in_h * up_y + pad_y0 + pad_y1 - kernel_h + 1, in_w * up_x + pad_x0 + pad_x1 - kernel_w + 1) + return out[:, :, ::down_y, ::down_x] + +def upfirdn2d(input, kernel, up=1, down=1, pad=(0, 0)): + return upfirdn2d_native(input, kernel, up, up, down, down, pad[0], pad[1], pad[0], pad[1]) + +# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/ops/fused_act/fused_act.py#L81 +class FusedLeakyReLU(torch.nn.Module): + def __init__(self, channel, negative_slope=0.2, scale=2 ** 0.5): + super().__init__() + self.bias = torch.nn.Parameter(torch.zeros(1, channel, 1, 1)) + self.negative_slope = negative_slope + self.scale = scale + + def forward(self, input): + return fused_leaky_relu(input, self.bias, self.negative_slope, self.scale) + +def fused_leaky_relu(input, bias, negative_slope=0.2, scale=2 ** 0.5): + return F.leaky_relu(input + bias, negative_slope) * scale + +class Blur(torch.nn.Module): + def __init__(self, kernel, pad): + super().__init__() + kernel = torch.tensor(kernel, dtype=torch.float32) + kernel = kernel[None, :] * kernel[:, None] + kernel = kernel / kernel.sum() + self.register_buffer('kernel', kernel) + self.pad = pad + + def forward(self, input): + return upfirdn2d(input, self.kernel, pad=self.pad) + +#https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L590 +class ScaledLeakyReLU(torch.nn.Module): + def __init__(self, negative_slope=0.2): + super().__init__() + self.negative_slope = negative_slope + + def forward(self, input): + return F.leaky_relu(input, negative_slope=self.negative_slope) + +# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L605 +class EqualConv2d(torch.nn.Module): + def __init__(self, in_channel, out_channel, kernel_size, stride=1, padding=0, bias=True): + super().__init__() + self.weight = torch.nn.Parameter(torch.randn(out_channel, in_channel, kernel_size, kernel_size)) + self.scale = 1 / math.sqrt(in_channel * kernel_size ** 2) + self.stride = stride + self.padding = padding + self.bias = torch.nn.Parameter(torch.zeros(out_channel)) if bias else None + + def forward(self, input): + return F.conv2d(input, self.weight * self.scale, bias=self.bias, stride=self.stride, padding=self.padding) + +# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L134 +class EqualLinear(torch.nn.Module): + def __init__(self, in_dim, out_dim, bias=True, bias_init=0, lr_mul=1, activation=None): + super().__init__() + self.weight = torch.nn.Parameter(torch.randn(out_dim, in_dim).div_(lr_mul)) + self.bias = torch.nn.Parameter(torch.zeros(out_dim).fill_(bias_init)) if bias else None + self.activation = activation + self.scale = (1 / math.sqrt(in_dim)) * lr_mul + self.lr_mul = lr_mul + + def forward(self, input): + if self.activation: + out = F.linear(input, self.weight * self.scale) + return fused_leaky_relu(out, self.bias * self.lr_mul) + return F.linear(input, self.weight * self.scale, bias=self.bias * self.lr_mul) + +# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L654 +class ConvLayer(torch.nn.Sequential): + def __init__(self, in_channel, out_channel, kernel_size, downsample=False, blur_kernel=[1, 3, 3, 1], bias=True, activate=True): + layers = [] + + if downsample: + factor = 2 + p = (len(blur_kernel) - factor) + (kernel_size - 1) + layers.append(Blur(blur_kernel, pad=((p + 1) // 2, p // 2))) + stride, padding = 2, 0 + else: + stride, padding = 1, kernel_size // 2 + + layers.append(EqualConv2d(in_channel, out_channel, kernel_size, padding=padding, stride=stride, bias=bias and not activate)) + + if activate: + layers.append(FusedLeakyReLU(out_channel) if bias else ScaledLeakyReLU(0.2)) + + super().__init__(*layers) + +# https://github.com/XPixelGroup/BasicSR/blob/8d56e3a045f9fb3e1d8872f92ee4a4f07f886b0a/basicsr/archs/stylegan2_arch.py#L704 +class ResBlock(torch.nn.Module): + def __init__(self, in_channel, out_channel): + super().__init__() + self.conv1 = ConvLayer(in_channel, in_channel, 3) + self.conv2 = ConvLayer(in_channel, out_channel, 3, downsample=True) + self.skip = ConvLayer(in_channel, out_channel, 1, downsample=True, activate=False, bias=False) + + def forward(self, input): + out = self.conv2(self.conv1(input)) + skip = self.skip(input) + return (out + skip) / math.sqrt(2) + + +class AppearanceEncoder(torch.nn.Module): + def __init__(self, w_dim=512): + super().__init__() + + self.convs = torch.nn.ModuleList([ + ConvLayer(3, 32, 1), ResBlock(32, 64), + ResBlock(64, 128), ResBlock(128, 256), + ResBlock(256, 512), ResBlock(512, 512), + ResBlock(512, 512), ResBlock(512, 512), + EqualConv2d(512, w_dim, 4, padding=0, bias=False) + ]) + + def forward(self, x): + for conv in self.convs: + x = conv(x) + return x.squeeze((-2, -1)) + +class MotionEncoder(torch.nn.Module): + def __init__(self, dim=512, motion_dim=20): + super().__init__() + self.net_app = AppearanceEncoder(dim) + self.fc = torch.nn.Sequential(*[EqualLinear(dim, dim) for _ in range(4)] + [EqualLinear(dim, motion_dim)]) + + def encode_motion(self, x): + return self.fc(self.net_app(x)) + +class MotionProjector(torch.nn.Module): + def __init__(self, m_dim): + super().__init__() + self.weight = torch.nn.Parameter(torch.randn(512, m_dim)) + self.motion_dim = m_dim + + def forward(self, input): + stabilized_weight = self.weight + 1e-8 * torch.eye(512, self.motion_dim, device=self.weight.device, dtype=self.weight.dtype) + Q, _ = torch.linalg.qr(stabilized_weight) + if input is None: + return Q + return torch.sum(input.unsqueeze(-1) * Q.T, dim=1) + +class MotionDecoder(torch.nn.Module): + def __init__(self, m_dim): + super().__init__() + self.direction = MotionProjector(m_dim) + +class MotionExtractor(torch.nn.Module): + def __init__(self, s_dim=512, m_dim=20): + super().__init__() + self.enc = MotionEncoder(s_dim, m_dim) + self.dec = MotionDecoder(m_dim) + + def forward(self, img): + motion_feat = self.enc.encode_motion(img) + return self.dec.direction(motion_feat) \ No newline at end of file