From 644f9aeaa3da436a78b063707ef88c523557190b Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 21 Aug 2025 22:14:37 +0300 Subject: [PATCH 01/16] Fix in multitalk sampling --- nodes.py | 1 + 1 file changed, 1 insertion(+) diff --git a/nodes.py b/nodes.py index 79d562e..d1635b2 100644 --- a/nodes.py +++ b/nodes.py @@ -2033,6 +2033,7 @@ class WanVideoSampler: context = get_context_scheduler(context_schedule) # vid2vid + noise_mask=None if samples is not None and not multitalk_sampling: saved_generator_state = samples.get("generator_state", None) if saved_generator_state is not None: From c2923068e6f0360e9cbb91de17b936afcdba7204 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 22 Aug 2025 01:24:16 +0300 Subject: [PATCH 02/16] Differential diffusion for Multi/InfiniteTalk long I2V as well --- nodes.py | 58 +++++++++++++++++++++++++++++--------------------------- 1 file changed, 30 insertions(+), 28 deletions(-) diff --git a/nodes.py b/nodes.py index d1635b2..678c8f4 100644 --- a/nodes.py +++ b/nodes.py @@ -2054,19 +2054,15 @@ class WanVideoSampler: original_image = input_samples.to(device) 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) 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 - ).squeeze(0) # Remove batch dim, keep channel dim - - # Add batch & channel dims for final output - noise_mask = noise_mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1) - - if noise_mask.shape[2] != noise.shape[1]: - noise_mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - noise_mask.shape[2], noise.shape[2], noise.shape[3]), noise_mask], dim=2) + ).repeat(1, noise.shape[0], 1, 1, 1) # extra latents (Pusa) and 5b latents_to_insert = add_index = None @@ -2626,16 +2622,13 @@ class WanVideoSampler: # diff diff prep masks = None if not multitalk_sampling and samples is not None and noise_mask is not None: - noise_mask = 1 - noise_mask thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps) - thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device) - masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device) - masks = masks > thresholds + thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device) + masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds latent_shift_loop = False if loop_args is not None: - latent_shift_loop = True - is_looped = True + latent_shift_loop = is_looped = True latent_skip = loop_args["shift_skip"] latent_shift_start_percent = loop_args["start_percent"] latent_shift_end_percent = loop_args["end_percent"] @@ -3125,6 +3118,8 @@ class WanVideoSampler: 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]}" @@ -3135,13 +3130,24 @@ class WanVideoSampler: noise = input_samples # diff diff prep - masks = None + noise_mask = samples.get("noise_mask", None) if noise_mask is not None: - noise_mask = 1 - noise_mask + 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.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device) - masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device) - masks = masks > thresholds + thresholds = thresholds.reshape(-1, 1, 1, 1, 1).to(device) + masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds window_vace_data = None if vace_data is not None: @@ -3292,10 +3298,10 @@ class WanVideoSampler: noise_pred, self.cache_state = predict_with_cfg( latent_model_input, - cfg[idx], + cfg[i], positive, text_embeds["negative_prompt_embeds"], - timestep, idx, y, clip_embeds, control_latents, window_vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, + timestep, i, y, clip_embeds, control_latents, window_vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, cache_state=self.cache_state, multitalk_audio_embeds=audio_embs) sampling_pbar.update(1) @@ -3309,8 +3315,7 @@ class WanVideoSampler: # update latent if scheduler == "multitalk": noise_pred = -noise_pred - dt = timesteps[i] - timesteps[i + 1] - dt = dt / 1000 + dt = (timesteps[i] - timesteps[i + 1]) / 1000 latent = latent + noise_pred * dt[:, None, None, None] else: latent = latent.to(intermediate_device) @@ -3323,12 +3328,9 @@ class WanVideoSampler: # differential diffusion inpaint if masks is not None: - if idx < len(timesteps) - 1: - noise_timestep = timesteps[idx+1] - image_latent = sample_scheduler.scale_noise( - original_image, torch.tensor([noise_timestep]), noise.to(device) - ) - mask = masks[idx].to(latent) + if i < len(timesteps) - 1: + image_latent = add_noise(original_image, noise.to(device), timesteps[i+1]) + mask = masks[i].to(latent) latent = image_latent * mask + latent * (1-mask) # injecting motion frames From bfee623d7713e5a331ff15e83c26ba133c3287cf Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 22 Aug 2025 01:48:35 +0300 Subject: [PATCH 03/16] Pass original image to 2nd sampler for 2.2 diff diff --- nodes.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/nodes.py b/nodes.py index 678c8f4..d610cea 100644 --- a/nodes.py +++ b/nodes.py @@ -2051,7 +2051,9 @@ class WanVideoSampler: noise_mask = samples.get("noise_mask", None) if noise_mask is not None: log.info(f"Latent noise_mask shape: {noise_mask.shape}") - original_image = input_samples.to(device) + original_image = samples.get("original_image", None) + if original_image is None: + original_image = input_samples if len(noise_mask.shape) == 4: noise_mask = noise_mask.squeeze(1) if noise_mask.shape[0] < noise.shape[1]: @@ -3329,7 +3331,7 @@ class WanVideoSampler: # differential diffusion inpaint if masks is not None: if i < len(timesteps) - 1: - image_latent = add_noise(original_image, noise.to(device), timesteps[i+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) @@ -3510,7 +3512,7 @@ class WanVideoSampler: if idx < len(timesteps) - 1: noise_timestep = timesteps[idx+1] image_latent = sample_scheduler.scale_noise( - original_image, torch.tensor([noise_timestep]), noise.to(device) + original_image.to(device), torch.tensor([noise_timestep]), noise.to(device) ) mask = masks[idx].to(latent) latent = image_latent * mask + latent * (1-mask) @@ -3565,6 +3567,7 @@ class WanVideoSampler: "has_ref": has_ref, "drop_last": drop_last, "generator_state": seed_g.get_state(), + "original_image": original_image.cpu() if original_image is not None else None },{ "samples": callback_latent.unsqueeze(0).cpu() if callback is not None else None, }) From 49951a85b6efae83455044e68eba5a9b7d4df331 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 22 Aug 2025 01:51:16 +0300 Subject: [PATCH 04/16] Update nodes.py --- nodes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/nodes.py b/nodes.py index d610cea..1e1d767 100644 --- a/nodes.py +++ b/nodes.py @@ -2033,7 +2033,7 @@ class WanVideoSampler: context = get_context_scheduler(context_schedule) # vid2vid - noise_mask=None + noise_mask=original_image=None if samples is not None and not multitalk_sampling: saved_generator_state = samples.get("generator_state", None) if saved_generator_state is not None: From ff1fe75919d360cb85d2d2a793c412720c6ee064 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 22 Aug 2025 12:48:12 +0300 Subject: [PATCH 05/16] RoPE ntk scaling --- nodes.py | 69 +++++++++++++++++++++++++++++---------- wanvideo/modules/model.py | 33 ++++++++++++++----- 2 files changed, 76 insertions(+), 26 deletions(-) diff --git a/nodes.py b/nodes.py index 1e1d767..94110aa 100644 --- a/nodes.py +++ b/nodes.py @@ -1528,7 +1528,37 @@ class WanVideoScheduler: #WIP def process(self, scheduler): return (scheduler,) - + +rope_functions = ["default", "comfy", "comfy_chunked"] +class WanVideoRoPEFunction: #WIP + @classmethod + def INPUT_TYPES(s): + return {"required": { + "rope_function": (rope_functions, {"default": "comfy"}), + "ntk_scale_f": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}), + "ntk_scale_h": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}), + "ntk_scale_w": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01}), + }, + } + + RETURN_TYPES = (rope_functions, ) + RETURN_NAMES = ("rope_function",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + EXPERIMENTAL = True + + def process(self, rope_function, ntk_scale_f, ntk_scale_h, ntk_scale_w): + if ntk_scale_f != 1.0 or ntk_scale_h != 1.0 or ntk_scale_w != 1.0: + rope_func_dict = { + "rope_function": rope_function, + "ntk_scale_f": ntk_scale_f, + "ntk_scale_h": ntk_scale_h, + "ntk_scale_w": ntk_scale_w, + } + return (rope_func_dict,) + return (rope_function,) + + #region Sampler class WanVideoSampler: @classmethod @@ -1555,7 +1585,7 @@ class WanVideoSampler: "flowedit_args": ("FLOWEDITARGS", ), "batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Batch cond and uncond for faster sampling, possibly faster on some hardware, uses more memory"}), "slg_args": ("SLGARGS", ), - "rope_function": (["default", "comfy", "comfy_chunked"], {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile. Chunked version has reduced peak VRAM usage when not using torch.compile"}), + "rope_function": (rope_functions, {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile. Chunked version has reduced peak VRAM usage when not using torch.compile"}), "loop_args": ("LOOPARGS", ), "experimental_args": ("EXPERIMENTALARGS", ), "sigmas": ("SIGMAS", ), @@ -1639,7 +1669,7 @@ class WanVideoSampler: log.info(f"sigmas: {sample_scheduler.sigmas}") else: timesteps = torch.tensor([1000, 750, 500, 250], device=device) - + total_steps = steps steps = len(timesteps) if end_step != -1 and start_step >= end_step: @@ -2043,8 +2073,8 @@ class WanVideoSampler: input_samples = torch.cat([input_samples[:, :1].repeat(1, noise.shape[1] - input_samples.shape[1], 1, 1), input_samples], dim=1) if add_noise_to_samples: - latent_timestep = timesteps[:1].to(noise) - noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples + latent_timestep = timesteps[:1].to(noise) + noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples else: noise = input_samples @@ -2110,10 +2140,7 @@ class WanVideoSampler: set_enhance_weight(feta_args["weight"]) feta_start_percent = feta_args["start_percent"] feta_end_percent = feta_args["end_percent"] - if context_options is not None: - set_num_frames(context_frames) - else: - set_num_frames(latent_video_length) + set_num_frames(latent_video_length) if context_options is None else set_num_frames(context_frames) enhance_enabled = True else: feta_args = None @@ -2267,6 +2294,11 @@ class WanVideoSampler: sample_scheduler_flipped = copy.deepcopy(sample_scheduler) #rope + ntk_alphas = [1.0, 1.0, 1.0] + if isinstance(rope_function, dict): + ntk_alphas = rope_function["ntk_scale_f"], rope_function["ntk_scale_h"], rope_function["ntk_scale_w"] + rope_function = rope_function["rope_function"] + freqs = None transformer.rope_embedder.k = None transformer.rope_embedder.num_frames = None @@ -2460,7 +2492,8 @@ class WanVideoSampler: "inner_t": [shot_len] if shot_len else None, "standin_input": standin_input, "fantasy_portrait_input": fantasy_portrait_input, - "reverse_time": reverse_time + "reverse_time": reverse_time, + "ntk_alphas": ntk_alphas } batch_size = 1 @@ -2576,13 +2609,13 @@ class WanVideoSampler: raise e #https://github.com/WeichenFan/CFG-Zero-star/ + alpha = 1.0 if use_cfg_zero_star: alpha = optimized_scale( noise_pred_cond.view(batch_size, -1), noise_pred_uncond.view(batch_size, -1) ).view(batch_size, 1, 1, 1) - else: - alpha = 1.0 + noise_pred_uncond_scaled = noise_pred_uncond * alpha @@ -2621,7 +2654,7 @@ class WanVideoSampler: intermediate_device = device - # diff diff prep + # Differential diffusion prep masks = None if not multitalk_sampling and samples is not None and noise_mask is not None: thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps) @@ -2676,10 +2709,8 @@ class WanVideoSampler: # Generate new random noise z_rand = torch.randn(z_T.shape, dtype=torch.float32, generator=seed_g, device=torch.device("cpu")) - # Apply frequency mixing - current_latent = freq_mix_3d(z_T.to(torch.float32), z_rand.to(device), LPF=freq_filter) - current_latent = current_latent.to(dtype) + current_latent = (freq_mix_3d(z_T.to(torch.float32), z_rand.to(device), LPF=freq_filter)).to(dtype) # Store initial noise for first iteration if freeinit_args is not None and iter_idx == 0: @@ -3798,7 +3829,8 @@ NODE_CLASS_MAPPINGS = { "WanVideoLatentReScale": WanVideoLatentReScale, "WanVideoScheduler": WanVideoScheduler, "WanVideoAddStandInLatent": WanVideoAddStandInLatent, - "WanVideoAddControlEmbeds": WanVideoAddControlEmbeds + "WanVideoAddControlEmbeds": WanVideoAddControlEmbeds, + "WanVideoRoPEFunction": WanVideoRoPEFunction } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3831,5 +3863,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoAddExtraLatent": "WanVideo Add Extra Latent", "WanVideoLatentReScale": "WanVideo Latent ReScale", "WanVideoAddStandInLatent": "WanVideo Add StandIn Latent", - "WanVideoAddControlEmbeds": "WanVideo Add Control Embeds" + "WanVideoAddControlEmbeds": "WanVideo Add Control Embeds", + "WanVideoRoPEFunction": "WanVideo RoPE Function" } diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 01da52c..e296894 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -96,18 +96,24 @@ def apply_rope_comfy_chunked(xq, xk, freqs_cis, num_chunks=4): return xq_out, xk_out -def rope_riflex(pos, dim, theta, L_test, k, temporal): +def rope_riflex(pos, dim, i, theta, L_test, k, ntk_factor=1.0): assert dim % 2 == 0 if mm.is_device_mps(pos.device) or mm.is_intel_xpu() or mm.is_directml_enabled(): device = torch.device("cpu") else: device = pos.device + if ntk_factor != 1.0: + print("scaling the theta with", ntk_factor) + + theta *= ntk_factor + print("theta", theta) + scale = torch.linspace(0, (dim - 2) / dim, steps=dim//2, dtype=torch.float64, device=device) omega = 1.0 / (theta**scale) # RIFLEX modification - adjust last frequency component if L_test and k are provided - if temporal and k > 0 and L_test: + if i==0 and k > 0 and L_test: omega[k-1] = 0.9 * 2 * torch.pi / L_test out = torch.einsum("...n,d->...nd", pos.to(dtype=torch.float32, device=device), omega) @@ -124,10 +130,18 @@ class EmbedND_RifleX(nn.Module): self.num_frames = num_frames self.k = k - def forward(self, ids): + def forward(self, ids, ntk_factor=[1.0,1.0,1.0]): n_axes = ids.shape[-1] emb = torch.cat( - [rope_riflex(ids[..., i], self.axes_dim[i], self.theta, self.num_frames, self.k, temporal=True if i == 0 else False) for i in range(n_axes)], + [rope_riflex( + ids[..., i], + self.axes_dim[i], + i, #f h w + self.theta, + self.num_frames, + self.k, + ntk_factor[i]) + for i in range(n_axes)], dim=-3, ) return emb.unsqueeze(1) @@ -1532,7 +1546,8 @@ class WanModel(torch.nn.Module): inner_t=None, standin_input=None, fantasy_portrait_input=None, - reverse_time=False + reverse_time=False, + ntk_alphas = [1.0, 1.0, 1.0] ): r""" Forward pass through the diffusion model @@ -1675,7 +1690,8 @@ class WanModel(torch.nn.Module): if (self.cached_freqs is not None and self.cached_shape == current_shape and self.cached_cond == has_cond and - self.cached_rope_k == self.rope_embedder.k + self.cached_rope_k == self.rope_embedder.k and + self.cached_ntk_alphas == ntk_alphas ): freqs = self.cached_freqs else: @@ -1706,15 +1722,16 @@ class WanModel(torch.nn.Module): combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1) # Generate RoPE frequencies for the combined positions - freqs = self.rope_embedder(combined_img_ids).movedim(1, 2) + freqs = self.rope_embedder(combined_img_ids, ntk_alphas).movedim(1, 2) else: img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1) - freqs = self.rope_embedder(img_ids).movedim(1, 2) + freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2) self.cached_freqs = freqs self.cached_shape = current_shape self.cached_cond = has_cond self.cached_rope_k = self.rope_embedder.k + self.cached_ntk_alphas = ntk_alphas # Stand-In RoPE frequencies if x_ip is not None: From bb5503c9bb2a6efe72c749f5104dd81715ca2a70 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 22 Aug 2025 12:57:57 +0300 Subject: [PATCH 06/16] remove prints --- wanvideo/modules/model.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index e296894..ab40673 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -104,10 +104,7 @@ def rope_riflex(pos, dim, i, theta, L_test, k, ntk_factor=1.0): device = pos.device if ntk_factor != 1.0: - print("scaling the theta with", ntk_factor) - theta *= ntk_factor - print("theta", theta) scale = torch.linspace(0, (dim - 2) / dim, steps=dim//2, dtype=torch.float64, device=device) omega = 1.0 / (theta**scale) From b1ceaaeb2e743db8881b9a44ba96e253e4bca389 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 22 Aug 2025 15:04:29 +0300 Subject: [PATCH 07/16] Support Wan22 VAE in WanVideoLatentReScale --- nodes.py | 68 +++++++++++++++++++++++++++++++----------- nodes_model_loading.py | 2 +- nodes_utility.py | 63 ++++++++++++++++++++++++++++++++++++-- 3 files changed, 113 insertions(+), 20 deletions(-) diff --git a/nodes.py b/nodes.py index 94110aa..6c8820b 100644 --- a/nodes.py +++ b/nodes.py @@ -1934,14 +1934,15 @@ class WanVideoSampler: dwpose_data = torch.cat([dwpose_data[:,:,:1].repeat(1,1,3,1,1), dwpose_data], dim=2) dwpose_data = transformer.dwpose_embedding(dwpose_data) log.info(f"UniAnimate pose embed shape: {dwpose_data.shape}") - if dwpose_data.shape[2] > latent_video_length: - log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is longer than the video length {latent_video_length}, truncating") - dwpose_data = dwpose_data[:,:, :latent_video_length] - elif dwpose_data.shape[2] < latent_video_length: - log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is shorter than the video length {latent_video_length}, padding with last pose") - pad_len = latent_video_length - dwpose_data.shape[2] - pad = dwpose_data[:,:,:1].repeat(1,1,pad_len,1,1) - dwpose_data = torch.cat([dwpose_data, pad], dim=2) + if not multitalk_sampling: + if dwpose_data.shape[2] > latent_video_length: + log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is longer than the video length {latent_video_length}, truncating") + dwpose_data = dwpose_data[:,:, :latent_video_length] + elif dwpose_data.shape[2] < latent_video_length: + log.warning(f"UniAnimate pose embed length {dwpose_data.shape[2]} is shorter than the video length {latent_video_length}, padding with last pose") + pad_len = latent_video_length - dwpose_data.shape[2] + pad = dwpose_data[:,:,:1].repeat(1,1,pad_len,1,1) + dwpose_data = torch.cat([dwpose_data, pad], dim=2) dwpose_data_flat = rearrange(dwpose_data, 'b c f h w -> b (f h w) c').contiguous() random_ref_dwpose_data = None @@ -3322,6 +3323,21 @@ class WanVideoSampler: else: positive = text_embeds["prompt_embeds"] + partial_unianim_data = None + if unianim_data is not None: + print(dwpose_data.shape) + partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx] + print("partial_dwpose shape:", partial_dwpose.shape) + partial_dwpose_flat=rearrange(partial_dwpose, 'b c f h w -> b (f h w) c') + print("partial_dwpose_flat shape:", partial_dwpose_flat.shape) + partial_unianim_data = { + "dwpose": partial_dwpose_flat, + "random_ref": unianim_data["random_ref"], + "strength": unianimate_poses["strength"], + "start_percent": unianimate_poses["start_percent"], + "end_percent": unianimate_poses["end_percent"] + } + sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True) for i in range(len(timesteps)-1): timestep = timesteps[i] @@ -3334,7 +3350,7 @@ class WanVideoSampler: cfg[i], positive, text_embeds["negative_prompt_embeds"], - timestep, i, y, clip_embeds, control_latents, window_vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, + timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond, cache_state=self.cache_state, multitalk_audio_embeds=audio_embs) sampling_pbar.update(1) @@ -3778,14 +3794,32 @@ class WanVideoLatentReScale: samples = samples.copy() latents = samples["samples"] - mean = [ - -0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508, - 0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921 - ] - std = [ - 2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743, - 3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160 - ] + if latents.shape[1] == 48: + mean = [ + -0.2289, -0.0052, -0.1323, -0.2339, -0.2799, 0.0174, 0.1838, 0.1557, + -0.1382, 0.0542, 0.2813, 0.0891, 0.1570, -0.0098, 0.0375, -0.1825, + -0.2246, -0.1207, -0.0698, 0.5109, 0.2665, -0.2108, -0.2158, 0.2502, + -0.2055, -0.0322, 0.1109, 0.1567, -0.0729, 0.0899, -0.2799, -0.1230, + -0.0313, -0.1649, 0.0117, 0.0723, -0.2839, -0.2083, -0.0520, 0.3748, + 0.0152, 0.1957, 0.1433, -0.2944, 0.3573, -0.0548, -0.1681, -0.0667, + ] + std = [ + 0.4765, 1.0364, 0.4514, 1.1677, 0.5313, 0.4990, 0.4818, 0.5013, + 0.8158, 1.0344, 0.5894, 1.0901, 0.6885, 0.6165, 0.8454, 0.4978, + 0.5759, 0.3523, 0.7135, 0.6804, 0.5833, 1.4146, 0.8986, 0.5659, + 0.7069, 0.5338, 0.4889, 0.4917, 0.4069, 0.4999, 0.6866, 0.4093, + 0.5709, 0.6065, 0.6415, 0.4944, 0.5726, 1.2042, 0.5458, 1.6887, + 0.3971, 1.0600, 0.3943, 0.5537, 0.5444, 0.4089, 0.7468, 0.7744 + ] + else: + mean = [ + -0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508, + 0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921 + ] + std = [ + 2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743, + 3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160 + ] mean = torch.tensor(mean).view(1, latents.shape[1], 1, 1, 1) std = torch.tensor(std).view(1, latents.shape[1], 1, 1, 1) inv_std = (1.0 / std).view(1, latents.shape[1], 1, 1, 1) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 4e1a6e2..4c91a8a 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1208,7 +1208,7 @@ class WanVideoModelLoader: desc=f"Loading transformer parameters to {transformer_load_device}", total=param_count, leave=True): - if "loras" in name: + if "loras" in name or "dwpose" in name or "randomref" in name: continue #print(name, param.dtype, param.device, param.shape) if isinstance(param, GGUFParameter): diff --git a/nodes_utility.py b/nodes_utility.py index 2f67b73..f6f5617 100644 --- a/nodes_utility.py +++ b/nodes_utility.py @@ -254,17 +254,76 @@ class DummyComfyWanModelObject: return None return (DummyModel(),) +class WanVideoLatentReScale: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "samples": ("LATENT",), + "direction": (["comfy_to_wrapper", "wrapper_to_comfy"], {"tooltip": "Direction to rescale latents, from comfy to wrapper or vice versa"}), + } + } + + RETURN_TYPES = ("LATENT",) + RETURN_NAMES = ("samples",) + FUNCTION = "encode" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Rescale latents to match the expected range for encoding or decoding. Can be used to " + + def encode(self, samples, direction): + samples = samples.copy() + latents = samples["samples"] + + if latents.shape[1] == 48: + mean = [ + -0.2289, -0.0052, -0.1323, -0.2339, -0.2799, 0.0174, 0.1838, 0.1557, + -0.1382, 0.0542, 0.2813, 0.0891, 0.1570, -0.0098, 0.0375, -0.1825, + -0.2246, -0.1207, -0.0698, 0.5109, 0.2665, -0.2108, -0.2158, 0.2502, + -0.2055, -0.0322, 0.1109, 0.1567, -0.0729, 0.0899, -0.2799, -0.1230, + -0.0313, -0.1649, 0.0117, 0.0723, -0.2839, -0.2083, -0.0520, 0.3748, + 0.0152, 0.1957, 0.1433, -0.2944, 0.3573, -0.0548, -0.1681, -0.0667, + ] + std = [ + 0.4765, 1.0364, 0.4514, 1.1677, 0.5313, 0.4990, 0.4818, 0.5013, + 0.8158, 1.0344, 0.5894, 1.0901, 0.6885, 0.6165, 0.8454, 0.4978, + 0.5759, 0.3523, 0.7135, 0.6804, 0.5833, 1.4146, 0.8986, 0.5659, + 0.7069, 0.5338, 0.4889, 0.4917, 0.4069, 0.4999, 0.6866, 0.4093, + 0.5709, 0.6065, 0.6415, 0.4944, 0.5726, 1.2042, 0.5458, 1.6887, + 0.3971, 1.0600, 0.3943, 0.5537, 0.5444, 0.4089, 0.7468, 0.7744 + ] + else: + mean = [ + -0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508, + 0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921 + ] + std = [ + 2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743, + 3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160 + ] + mean = torch.tensor(mean).view(1, latents.shape[1], 1, 1, 1) + std = torch.tensor(std).view(1, latents.shape[1], 1, 1, 1) + inv_std = (1.0 / std).view(1, latents.shape[1], 1, 1, 1) + if direction == "comfy_to_wrapper": + latents = (latents - mean.to(latents)) * inv_std.to(latents) + elif direction == "wrapper_to_comfy": + latents = latents / inv_std.to(latents) + mean.to(latents) + + samples["samples"] = latents + + return (samples,) + NODE_CLASS_MAPPINGS = { "WanVideoImageResizeToClosest": WanVideoImageResizeToClosest, "WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame, "ExtractStartFramesForContinuations": ExtractStartFramesForContinuations, "CreateCFGScheduleFloatList": CreateCFGScheduleFloatList, - "DummyComfyWanModelObject": DummyComfyWanModelObject + "DummyComfyWanModelObject": DummyComfyWanModelObject, + "WanVideoLatentReScale": WanVideoLatentReScale } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest", "WanVideoVACEStartToEndFrame": "WanVideo VACE Start To End Frame", "ExtractStartFramesForContinuations": "Extract Start Frames For Continuations", "CreateCFGScheduleFloatList": "Create CFG Schedule Float List", - "DummyComfyWanModelObject": "Dummy Comfy Wan Model Object" + "DummyComfyWanModelObject": "Dummy Comfy Wan Model Object", + "WanVideoLatentReScale": "WanVideo Latent ReScale" } \ No newline at end of file From 948fdd369d5266bb8023bbebb26d79efbca66420 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 01:05:41 +0300 Subject: [PATCH 08/16] cleanup --- nodes.py | 59 ------------------------------------------------ nodes_utility.py | 2 +- 2 files changed, 1 insertion(+), 60 deletions(-) diff --git a/nodes.py b/nodes.py index 6c8820b..48716f3 100644 --- a/nodes.py +++ b/nodes.py @@ -3775,63 +3775,6 @@ class WanVideoEncode: return ({"samples": latents, "noise_mask": mask},) -class WanVideoLatentReScale: - @classmethod - def INPUT_TYPES(s): - return {"required": { - "samples": ("LATENT",), - "direction": (["comfy_to_wrapper", "wrapper_to_comfy"], {"tooltip": "Direction to rescale latents, from comfy to wrapper or vice versa"}), - } - } - - RETURN_TYPES = ("LATENT",) - RETURN_NAMES = ("samples",) - FUNCTION = "encode" - CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Rescale latents to match the expected range for encoding or decoding. Can be used to " - - def encode(self, samples, direction): - samples = samples.copy() - latents = samples["samples"] - - if latents.shape[1] == 48: - mean = [ - -0.2289, -0.0052, -0.1323, -0.2339, -0.2799, 0.0174, 0.1838, 0.1557, - -0.1382, 0.0542, 0.2813, 0.0891, 0.1570, -0.0098, 0.0375, -0.1825, - -0.2246, -0.1207, -0.0698, 0.5109, 0.2665, -0.2108, -0.2158, 0.2502, - -0.2055, -0.0322, 0.1109, 0.1567, -0.0729, 0.0899, -0.2799, -0.1230, - -0.0313, -0.1649, 0.0117, 0.0723, -0.2839, -0.2083, -0.0520, 0.3748, - 0.0152, 0.1957, 0.1433, -0.2944, 0.3573, -0.0548, -0.1681, -0.0667, - ] - std = [ - 0.4765, 1.0364, 0.4514, 1.1677, 0.5313, 0.4990, 0.4818, 0.5013, - 0.8158, 1.0344, 0.5894, 1.0901, 0.6885, 0.6165, 0.8454, 0.4978, - 0.5759, 0.3523, 0.7135, 0.6804, 0.5833, 1.4146, 0.8986, 0.5659, - 0.7069, 0.5338, 0.4889, 0.4917, 0.4069, 0.4999, 0.6866, 0.4093, - 0.5709, 0.6065, 0.6415, 0.4944, 0.5726, 1.2042, 0.5458, 1.6887, - 0.3971, 1.0600, 0.3943, 0.5537, 0.5444, 0.4089, 0.7468, 0.7744 - ] - else: - mean = [ - -0.7571, -0.7089, -0.9113, 0.1075, -0.1745, 0.9653, -0.1517, 1.5508, - 0.4134, -0.0715, 0.5517, -0.3632, -0.1922, -0.9497, 0.2503, -0.2921 - ] - std = [ - 2.8184, 1.4541, 2.3275, 2.6558, 1.2196, 1.7708, 2.6052, 2.0743, - 3.2687, 2.1526, 2.8652, 1.5579, 1.6382, 1.1253, 2.8251, 1.9160 - ] - mean = torch.tensor(mean).view(1, latents.shape[1], 1, 1, 1) - std = torch.tensor(std).view(1, latents.shape[1], 1, 1, 1) - inv_std = (1.0 / std).view(1, latents.shape[1], 1, 1, 1) - if direction == "comfy_to_wrapper": - latents = (latents - mean.to(latents)) * inv_std.to(latents) - elif direction == "wrapper_to_comfy": - latents = latents / inv_std.to(latents) + mean.to(latents) - - samples["samples"] = latents - - return (samples,) - NODE_CLASS_MAPPINGS = { "WanVideoSampler": WanVideoSampler, "WanVideoDecode": WanVideoDecode, @@ -3860,7 +3803,6 @@ NODE_CLASS_MAPPINGS = { "WanVideoBlockList": WanVideoBlockList, "WanVideoTextEncodeCached": WanVideoTextEncodeCached, "WanVideoAddExtraLatent": WanVideoAddExtraLatent, - "WanVideoLatentReScale": WanVideoLatentReScale, "WanVideoScheduler": WanVideoScheduler, "WanVideoAddStandInLatent": WanVideoAddStandInLatent, "WanVideoAddControlEmbeds": WanVideoAddControlEmbeds, @@ -3895,7 +3837,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoBlockList": "WanVideo Block List", "WanVideoTextEncodeCached": "WanVideo TextEncode Cached", "WanVideoAddExtraLatent": "WanVideo Add Extra Latent", - "WanVideoLatentReScale": "WanVideo Latent ReScale", "WanVideoAddStandInLatent": "WanVideo Add StandIn Latent", "WanVideoAddControlEmbeds": "WanVideo Add Control Embeds", "WanVideoRoPEFunction": "WanVideo RoPE Function" diff --git a/nodes_utility.py b/nodes_utility.py index f6f5617..23c1d22 100644 --- a/nodes_utility.py +++ b/nodes_utility.py @@ -267,7 +267,7 @@ class WanVideoLatentReScale: RETURN_NAMES = ("samples",) FUNCTION = "encode" CATEGORY = "WanVideoWrapper" - DESCRIPTION = "Rescale latents to match the expected range for encoding or decoding. Can be used to " + DESCRIPTION = "Rescale latents to match the expected range for encoding or decoding between native ComfyUI VAE and the WanVideoWrapper VAE." def encode(self, samples, direction): samples = samples.copy() From d0b9f2c90781d8b0e2078ced29a4f0a6d25fd69d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 01:06:22 +0300 Subject: [PATCH 09/16] Allow UniAnimate to work with unmerged LoRAs --- nodes_model_loading.py | 12 ++++++++++-- unianimate/nodes.py | 13 +++++++------ 2 files changed, 17 insertions(+), 8 deletions(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 4c91a8a..f3825ed 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1106,6 +1106,7 @@ class WanVideoModelLoader: patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) patcher.model.is_patched = False + unianimate_sd = None control_lora = False if lora is not None: for l in lora: @@ -1123,7 +1124,7 @@ class WanVideoModelLoader: if "dwpose_embedding.0.weight" in lora_sd: #unianimate from .unianimate.nodes import update_transformer log.info("Unianimate LoRA detected, patching model...") - transformer = update_transformer(transformer, lora_sd) + transformer, unianimate_sd = update_transformer(transformer, lora_sd) lora_sd = standardize_lora_key_format(lora_sd) @@ -1192,6 +1193,13 @@ class WanVideoModelLoader: low_mem_load=lora_low_mem_load, control_lora=control_lora, scale_weights=scale_weights) scale_weights.clear() patcher.patches.clear() + + if unianimate_sd is not None: + sd.update(unianimate_sd) + for name, param in transformer.named_parameters(): + if "dwpose_embedding" in name or "randomref_embedding_pose" in name: + dtype_to_use = base_dtype + set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) if gguf: #from diffusers.quantizers.gguf.utils import _replace_with_gguf_linear, GGUFParameter @@ -1208,7 +1216,7 @@ class WanVideoModelLoader: desc=f"Loading transformer parameters to {transformer_load_device}", total=param_count, leave=True): - if "loras" in name or "dwpose" in name or "randomref" in name: + if "loras" in name: continue #print(name, param.dtype, param.device, param.shape) if isinstance(param, GGUFParameter): diff --git a/unianimate/nodes.py b/unianimate/nodes.py index c533a2a..eb54de4 100644 --- a/unianimate/nodes.py +++ b/unianimate/nodes.py @@ -37,17 +37,18 @@ def update_transformer(transformer, state_dict): nn.SiLU(), nn.Conv2d(concat_dim * 4, randomref_dim, 3, stride=2, padding=1), ) + unianimate_sd = {} state_dict_new = {} for key in list(state_dict.keys()): if "dwpose_embedding" in key: - state_dict_new[key.split("dwpose_embedding.")[1]] = state_dict.pop(key) - transformer.dwpose_embedding.load_state_dict(state_dict_new, strict=True) - state_dict_new = {} + state_dict_new[key] = state_dict.pop(key) + unianimate_sd.update(state_dict_new) for key in list(state_dict.keys()): if "randomref_embedding_pose" in key: - state_dict_new[key.split("randomref_embedding_pose.")[1]] = state_dict.pop(key) - transformer.randomref_embedding_pose.load_state_dict(state_dict_new,strict=True) - return transformer + state_dict_new[key] = state_dict.pop(key) + unianimate_sd.update(state_dict_new) + del state_dict_new + return transformer, unianimate_sd # Openpose # Original from CMU https://github.com/CMU-Perceptual-Computing-Lab/openpose From 4eeaf1ea194ed32e0a2fef2a16201c450a8da40f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 01:06:57 +0300 Subject: [PATCH 10/16] Update nodes.py --- nodes.py | 3 --- 1 file changed, 3 deletions(-) diff --git a/nodes.py b/nodes.py index 48716f3..9a51b14 100644 --- a/nodes.py +++ b/nodes.py @@ -3325,11 +3325,8 @@ class WanVideoSampler: partial_unianim_data = None if unianim_data is not None: - print(dwpose_data.shape) partial_dwpose = dwpose_data[:, :, latent_start_idx:latent_end_idx] - print("partial_dwpose shape:", partial_dwpose.shape) partial_dwpose_flat=rearrange(partial_dwpose, 'b c f h w -> b (f h w) c') - print("partial_dwpose_flat shape:", partial_dwpose_flat.shape) partial_unianim_data = { "dwpose": partial_dwpose_flat, "random_ref": unianim_data["random_ref"], From 5f521dc1692a7f2ac872514b092a991ab78fedd5 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 13:04:19 +0300 Subject: [PATCH 11/16] Add generic CreateScheduleFloatList to assist with LoRA scheduling --- nodes_utility.py | 92 +++++++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 87 insertions(+), 5 deletions(-) diff --git a/nodes_utility.py b/nodes_utility.py index 23c1d22..3cded7b 100644 --- a/nodes_utility.py +++ b/nodes_utility.py @@ -2,6 +2,11 @@ import torch import numpy as np from comfy.utils import common_upscale +try: + from server import PromptServer +except: + PromptServer = None + VAE_STRIDE = (4, 8, 8) PATCH_SIZE = (1, 2, 2) @@ -188,7 +193,10 @@ class CreateCFGScheduleFloatList: "interpolation": (["linear", "ease_in", "ease_out"], {"default": "linear", "tooltip": "Interpolation method to use for the cfg scale"}), "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "Start percent of the steps to apply cfg"}), "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "End percent of the steps to apply cfg"}), - } + }, + "hidden": { + "unique_id": "UNIQUE_ID", + }, } RETURN_TYPES = ("FLOAT", ) @@ -197,8 +205,8 @@ class CreateCFGScheduleFloatList: CATEGORY = "WanVideoWrapper" DESCRIPTION = "Helper node to generate a list of floats that can be used to schedule cfg scale for the steps, outside the set range cfg is set to 1.0" - def process(self, steps, cfg_scale_start, cfg_scale_end, interpolation, start_percent, end_percent): - + def process(self, steps, cfg_scale_start, cfg_scale_end, interpolation, start_percent, end_percent, unique_id): + # Create a list of floats for the cfg schedule cfg_list = [1.0] * steps start_idx = min(int(steps * start_percent), steps - 1) @@ -226,6 +234,78 @@ class CreateCFGScheduleFloatList: if start_percent > 0: cfg_list[0] = 1.0 + if unique_id and PromptServer is not None: + try: + PromptServer.instance.send_progress_text( + f"{cfg_list}", + unique_id + ) + except: + pass + + return (cfg_list,) + +class CreateScheduleFloatList: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "steps": ("INT", {"default": 30, "min": 2, "max": 1000, "step": 1, "tooltip": "Number of steps to schedule cfg for"} ), + "start_value": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}), + "end_value": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 100.0, "step": 0.01, "round": 0.01, "tooltip": "CFG scale to use for the steps"}), + "default_value": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1000.0, "step": 0.01, "round": 0.01, "tooltip": "Default value to use for the steps"}), + "interpolation": (["linear", "ease_in", "ease_out"], {"default": "linear", "tooltip": "Interpolation method to use for the cfg scale"}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "Start percent of the steps to apply cfg"}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.01,"tooltip": "End percent of the steps to apply cfg"}), + }, + "hidden": { + "unique_id": "UNIQUE_ID", + }, + } + + RETURN_TYPES = ("FLOAT", ) + RETURN_NAMES = ("float_list",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Helper node to generate a list of floats that can be used to schedule things like cfg and lora scale per step" + + def process(self, steps, start_value, end_value, default_value,interpolation, start_percent, end_percent, unique_id): + + # Create a list of floats for the cfg schedule + cfg_list = [default_value] * steps + start_idx = min(int(steps * start_percent), steps - 1) + end_idx = min(int(steps * end_percent), steps - 1) + + for i in range(start_idx, end_idx + 1): + if i >= steps: + break + + if end_idx == start_idx: + t = 0 + else: + t = (i - start_idx) / (end_idx - start_idx) + + if interpolation == "linear": + factor = t + elif interpolation == "ease_in": + factor = t * t + elif interpolation == "ease_out": + factor = t * (2 - t) + + cfg_list[i] = round(start_value + factor * (end_value - start_value), 2) + + # If start_percent > 0, always include the first step + if start_percent > 0: + cfg_list[0] = default_value + + if unique_id and PromptServer is not None: + try: + PromptServer.instance.send_progress_text( + f"{cfg_list}", + unique_id + ) + except: + pass + return (cfg_list,) @@ -317,7 +397,8 @@ NODE_CLASS_MAPPINGS = { "ExtractStartFramesForContinuations": ExtractStartFramesForContinuations, "CreateCFGScheduleFloatList": CreateCFGScheduleFloatList, "DummyComfyWanModelObject": DummyComfyWanModelObject, - "WanVideoLatentReScale": WanVideoLatentReScale + "WanVideoLatentReScale": WanVideoLatentReScale, + "CreateScheduleFloatList": CreateScheduleFloatList } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest", @@ -325,5 +406,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ExtractStartFramesForContinuations": "Extract Start Frames For Continuations", "CreateCFGScheduleFloatList": "Create CFG Schedule Float List", "DummyComfyWanModelObject": "Dummy Comfy Wan Model Object", - "WanVideoLatentReScale": "WanVideo Latent ReScale" + "WanVideoLatentReScale": "WanVideo Latent ReScale", + "CreateScheduleFloatList": "Create Schedule Float List" } \ No newline at end of file From 73cff6ebab34c4a5e0043fe767c48885db607447 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 16:27:10 +0300 Subject: [PATCH 12/16] Cleanup Multi/InfiniteTalk sampling loop some --- nodes.py | 85 +++++++++++++++++++++----------------------------------- 1 file changed, 32 insertions(+), 53 deletions(-) diff --git a/nodes.py b/nodes.py index 9a51b14..7a10db2 100644 --- a/nodes.py +++ b/nodes.py @@ -3085,7 +3085,6 @@ class WanVideoSampler: } estimated_iterations = total_frames // (frame_num - motion_frame) + 1 - loop_pbar = tqdm(total=estimated_iterations, desc="Total progress", position=1, leave=True) callback = prepare_callback(patcher, estimated_iterations) audio_embedding = multitalk_audio_embedding @@ -3221,15 +3220,11 @@ class WanVideoSampler: # encode vae.to(device) y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype) - if mode == "multitalk": latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W else: - if is_first_clip: - latent_motion_frames = vae.encode(cond_image.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype) - else: - latent_motion_frames = vae.encode(cond_frame.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype) - latent_motion_frames = latent_motion_frames[0] + cond_ = cond_image if is_first_clip else cond_frame + latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0] vae.to(offload_device) y = torch.concat([msk, y], dim=1).squeeze(0) # 4+C T H W mm.soft_empty_cache() @@ -3279,8 +3274,7 @@ class WanVideoSampler: latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device) motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous() add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0]) - _, T_m, _, _ = add_latent.shape - latent[:, :T_m] = add_latent + latent[:, :add_latent.shape[1]] = add_latent if offload: #blockswap init @@ -3343,19 +3337,16 @@ class WanVideoSampler: latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames noise_pred, self.cache_state = predict_with_cfg( - latent_model_input, - cfg[i], - positive, - text_embeds["negative_prompt_embeds"], + latent_model_input, cfg[i], positive, text_embeds["negative_prompt_embeds"], timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond, cache_state=self.cache_state, multitalk_audio_embeds=audio_embs) - sampling_pbar.update(1) - 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)-1)) + del callback_latent + sampling_pbar.update(1) step_iteration_count += 1 # update latent @@ -3364,13 +3355,8 @@ class WanVideoSampler: dt = (timesteps[i] - timesteps[i + 1]) / 1000 latent = latent + noise_pred * dt[:, None, None, None] else: - latent = latent.to(intermediate_device) - temp_x0 = sample_scheduler.step( - noise_pred.unsqueeze(0), - timestep, - latent.unsqueeze(0), - **scheduler_step_args)[0] - latent = temp_x0.squeeze(0) + 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: @@ -3384,62 +3370,56 @@ class WanVideoSampler: latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device) motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous() add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1]) - _, T_m, _, _ = add_latent.shape - latent[:, :T_m] = add_latent + latent[:, :add_latent.shape[1]] = add_latent else: latent[:, :cur_motion_frames_latent_num] = latent_motion_frames - x0 = latent.to(device) - del latent_model_input, timestep - + del noise, y, msk, latent_motion_frames if offload: transformer.to(offload_device) vae.to(device) - videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae, pbar=False) + videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu() vae.to(offload_device) sampling_pbar.close() - # cache generated samples - videos = torch.stack(videos).cpu() # B C T H W + # optional color correction (less relevant for InfiniteTalk) if colormatch != "disabled": - videos = videos[0].permute(1, 2, 3, 0).cpu().float().numpy() + videos = videos.permute(1, 2, 3, 0).float().numpy() from color_matcher import ColorMatcher cm = ColorMatcher() cm_result_list = [] for img in videos: if mode == "multitalk": - cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch) + cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch) else: - cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch) - cm_result_list.append(torch.from_numpy(cm_result)) + cm_result = cm.transfer(src=img, ref=cond_image[0].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).to(torch.float32).permute(3, 0, 1, 2).unsqueeze(0) + videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2) + + # cache generated samples + gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:]) - if is_first_clip: - gen_video_list.append(videos) - else: - gen_video_list.append(videos[:, :, cur_motion_frames_num:]) current_condframe_index += 1 + iteration_count += 1 # decide whether is done if arrive_last_frame: - loop_pbar.update(estimated_iterations - iteration_count) - loop_pbar.close() break # update next condition frames is_first_clip = False cur_motion_frames_num = motion_frame + cond_ = videos[:, -cur_motion_frames_num:].unsqueeze(0) if mode == "infinitetalk": - cond_frame = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device) + cond_frame = cond_ else: - cond_image = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device) + cond_image = cond_ - # Update progress bar - iteration_count += 1 - loop_pbar.update(1) + del videos, latent + mm.soft_empty_cache() # Repeat audio emb if multitalk_embeds is not None: @@ -3464,9 +3444,8 @@ class WanVideoSampler: miss_length = 1 original_images = torch.cat([original_images, last_frame.repeat(1, 1, miss_length, 1, 1)], dim=2) - gen_video_samples = torch.cat(gen_video_list, dim=2).to(torch.float32) - - del noise, latent + gen_video_samples = torch.cat(gen_video_list, dim=1) + if force_offload: if not model["auto_cpu_offload"]: offload_transformer(transformer) @@ -3475,7 +3454,7 @@ class WanVideoSampler: torch.cuda.reset_peak_memory_stats(device) except: pass - return {"video": gen_video_samples[0].permute(1, 2, 3, 0).cpu()}, + return {"video": gen_video_samples.permute(1, 2, 3, 0)}, #region normal inference else: @@ -3657,9 +3636,9 @@ class WanVideoDecode: mm.soft_empty_cache() video = samples.get("video", None) if video is not None: - video = torch.clamp(video, -1.0, 1.0) - video = (video + 1.0) / 2.0 - return video.cpu(), + video.clamp_(-1.0, 1.0) + video.add_(1.0).div_(2.0) + return video.cpu().float(), latents = samples["samples"] end_image = samples.get("end_image", None) has_ref = samples.get("has_ref", False) From 9a2a13498a92a5547863ded8591a35fb82de8636 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 17:49:04 +0300 Subject: [PATCH 13/16] Cleanup MultiTalk model code --- multitalk/multitalk.py | 104 ++++++++++------------------------------- multitalk/nodes.py | 1 - nodes_model_loading.py | 5 +- 3 files changed, 25 insertions(+), 85 deletions(-) diff --git a/multitalk/multitalk.py b/multitalk/multitalk.py index e8afb3f..cdec103 100644 --- a/multitalk/multitalk.py +++ b/multitalk/multitalk.py @@ -1,11 +1,8 @@ -from diffusers import ModelMixin, ConfigMixin from einops import rearrange, repeat import torch import torch.nn as nn from ..wanvideo.modules.attention import attention -from comfy import model_management as mm - def timestep_transform( t, shift=5.0, @@ -65,7 +62,6 @@ def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', att x_ref_attn_map_source = x_ref_attn_map_source.to(visual_q.dtype) for class_idx, ref_target_mask in enumerate(ref_target_masks): - mm.soft_empty_cache() ref_target_mask = ref_target_mask[None, None, None, ...] x_ref_attnmap = x_ref_attn_map_source * ref_target_mask x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens @@ -78,13 +74,11 @@ def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', att x_ref_attn_maps.append(x_ref_attnmap) - del attn - del x_ref_attn_map_source - mm.soft_empty_cache() + del attn, x_ref_attn_map_source return torch.concat(x_ref_attn_maps, dim=0) -def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2, enable_sp=False): +def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2): """Args: query (torch.tensor): B M H K key (torch.tensor): B M H K @@ -145,7 +139,7 @@ class RotaryPositionalEmbedding1D(nn.Module): return x_.type_as(x) -class AudioProjModel(ModelMixin, ConfigMixin): +class AudioProjModel(nn.Module): def __init__( self, seq_len=5, @@ -217,11 +211,6 @@ class SingleStreamAttention(nn.Module): encoder_hidden_states_dim: int, num_heads: int, qkv_bias: bool, - qk_norm: bool, - norm_layer: nn.Module, - attn_drop: float = 0.0, - proj_drop: float = 0.0, - eps: float = 1e-6, attention_mode: str = 'sdpa', ) -> None: super().__init__() @@ -230,76 +219,46 @@ class SingleStreamAttention(nn.Module): self.encoder_hidden_states_dim = encoder_hidden_states_dim self.num_heads = num_heads self.head_dim = dim // num_heads - self.scale = self.head_dim**-0.5 - self.qk_norm = qk_norm - - self.q_linear = nn.Linear(dim, dim, bias=qkv_bias) - - self.q_norm = norm_layer(self.head_dim, eps=eps) if qk_norm else nn.Identity() - self.k_norm = norm_layer(self.head_dim,eps=eps) if qk_norm else nn.Identity() - - self.attn_drop = nn.Dropout(attn_drop) - self.proj = nn.Linear(dim, dim) - self.proj_drop = nn.Dropout(proj_drop) - - self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias) - - self.add_q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() - self.add_k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity() - self.attention_mode = attention_mode - def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor: + self.q_linear = nn.Linear(dim, dim, bias=qkv_bias) + self.proj = nn.Linear(dim, dim) + self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias) + + def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None) -> torch.Tensor: N_t, N_h, N_w = shape + expected_tokens = N_t * N_h * N_w + actual_tokens = x.shape[1] x_extra = None - try: - x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) - except: + + if actual_tokens != expected_tokens: x_extra = x[:, -N_h * N_w:, :] x = x[:, :-N_h * N_w, :] N_t = N_t - 1 - x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t) + + B = x.shape[0] + S = N_h * N_w + x = x.view(B * N_t, S, self.dim) # get q for hidden_state - B, N, C = x.shape - q = self.q_linear(x) - q_shape = (B, N, self.num_heads, self.head_dim) - q = q.view(q_shape).permute((0, 2, 1, 3)) - - if self.qk_norm: - q = self.q_norm(q) + q = self.q_linear(x).view(B * N_t, S, self.num_heads, self.head_dim) - # get kv from encoder_hidden_states - _, N_a, _ = encoder_hidden_states.shape - encoder_kv = self.kv_linear(encoder_hidden_states) - encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim) - encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4)) - encoder_k, encoder_v = encoder_kv.unbind(0) + # get kv from encoder_hidden_states # shape: (B, N, num_heads, head_dim) + kv = self.kv_linear(encoder_hidden_states) + encoder_k, encoder_v = kv.view(B * N_t, encoder_hidden_states.shape[1], 2, self.num_heads, self.head_dim).unbind(2) - if self.qk_norm: - encoder_k = self.add_k_norm(encoder_k) - - x = attention( - q.transpose(1, 2), - encoder_k.transpose(1, 2), - encoder_v.transpose(1, 2), - attention_mode=self.attention_mode - ) + x = attention(q, encoder_k, encoder_v, attention_mode=self.attention_mode) # linear transform - x_output_shape = (B, N, C) - #x = x.transpose(1, 2) - x = x.reshape(x_output_shape) - x = self.proj(x) - x = self.proj_drop(x) - - x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) + x = self.proj(x.reshape(B * N_t, S, self.dim)) + x = x.view(B, N_t * S, self.dim) if x_extra is not None: x = torch.cat([x, torch.zeros_like(x_extra)], dim=1) return x + class SingleStreamMultiAttention(SingleStreamAttention): """Multi-speaker rotary-position cross-attention. @@ -317,11 +276,6 @@ class SingleStreamMultiAttention(SingleStreamAttention): encoder_hidden_states_dim: int, num_heads: int, qkv_bias: bool, - qk_norm: bool, - norm_layer: nn.Module, - attn_drop: float = 0.0, - proj_drop: float = 0.0, - eps: float = 1e-6, class_range: int = 24, class_interval: int = 4, attention_mode: str = 'sdpa', @@ -331,11 +285,6 @@ class SingleStreamMultiAttention(SingleStreamAttention): encoder_hidden_states_dim=encoder_hidden_states_dim, num_heads=num_heads, qkv_bias=qkv_bias, - qk_norm=qk_norm, - norm_layer=norm_layer, - attn_drop=attn_drop, - proj_drop=proj_drop, - eps=eps, attention_mode=attention_mode, ) @@ -378,8 +327,6 @@ class SingleStreamMultiAttention(SingleStreamAttention): B, N, C = x.shape q = self.q_linear(x) q = q.view(B, N, self.num_heads, self.head_dim).permute(0, 2, 1, 3) - if self.qk_norm: - q = self.q_norm(q) if human_num == 2: # Use `class_range` logic for exactly 2 speakers @@ -443,8 +390,6 @@ class SingleStreamMultiAttention(SingleStreamAttention): encoder_kv = self.kv_linear(encoder_hidden_states) encoder_kv = encoder_kv.view(B, N_a, 2, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) encoder_k, encoder_v = encoder_kv.unbind(0) - if self.qk_norm: - encoder_k = self.add_k_norm(encoder_k) # Rotary for keys – assign centre of each speaker bucket to its context tokens if human_num == 2: @@ -480,7 +425,6 @@ class SingleStreamMultiAttention(SingleStreamAttention): # Linear projection x = x.reshape(B, N, C) x = self.proj(x) - x = self.proj_drop(x) # Restore original layout x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t) diff --git a/multitalk/nodes.py b/multitalk/nodes.py index 628c5d2..cfd7428 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -2,7 +2,6 @@ import folder_paths from comfy import model_management as mm from comfy.utils import load_torch_file, common_upscale from accelerate import init_empty_weights -from accelerate.utils import set_module_tensor_to_device import torch from ..utils import log diff --git a/nodes_model_loading.py b/nodes_model_loading.py index f3825ed..76fb7b4 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1017,7 +1017,7 @@ class WanVideoModelLoader: multitalk_model_type = multitalk_model.get("model_type", "MultiTalk") # init audio module from .multitalk.multitalk import SingleStreamMultiAttention - from .wanvideo.modules.model import WanRMSNorm, WanLayerNorm + from .wanvideo.modules.model import WanLayerNorm norm_input_visual = True #dunno what this is for block in transformer.blocks: @@ -1025,10 +1025,7 @@ class WanVideoModelLoader: dim=dim, encoder_hidden_states_dim=768, num_heads=num_heads, - qk_norm=False, qkv_bias=True, - eps=transformer.eps, - norm_layer=WanRMSNorm, class_range=24, class_interval=4, attention_mode=attention_mode, From 77d79f68781dfed1865435416268360385c6951e Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 17:55:12 +0300 Subject: [PATCH 14/16] This is redundant --- nodes.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/nodes.py b/nodes.py index 7a10db2..54b2f94 100644 --- a/nodes.py +++ b/nodes.py @@ -1681,8 +1681,6 @@ class WanVideoSampler: start_step = steps - int(steps * denoise_strength) - 1 add_noise_to_samples = True #for now to not break old workflows - first_sampler = (end_step != -1 or end_step >= steps) - noise_pred_flipped = None if isinstance(cfg, list): @@ -1692,7 +1690,7 @@ class WanVideoSampler: else: cfg = [cfg] * (steps + 1) - if first_sampler: + if end_step != -1: timesteps = timesteps[:end_step] sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1] log.info(f"Sampling until step {end_step}, timestep: {timesteps[-1]}") @@ -3247,7 +3245,7 @@ class WanVideoSampler: if start_step != 0: raise ValueError("start_step must be 0 when denoise_strength is used") start_step = steps - int(steps * denoise_strength) - 1 - if (end_step != -1 or end_step >= steps): + if end_step != -1: timesteps = timesteps[:end_step] sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1] if start_step > 0: From b7d7f9afe5b0cab01b7214e6191aee2a1039f186 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 18:59:21 +0300 Subject: [PATCH 15/16] cleanup unianimate stuff --- unianimate/dwpose/onnxdet.py | 127 ------------ unianimate/dwpose/onnxpose.py | 360 --------------------------------- unianimate/dwpose/util.py | 16 +- unianimate/dwpose/wholebody.py | 1 - unianimate/nodes.py | 21 +- 5 files changed, 17 insertions(+), 508 deletions(-) delete mode 100644 unianimate/dwpose/onnxdet.py delete mode 100644 unianimate/dwpose/onnxpose.py diff --git a/unianimate/dwpose/onnxdet.py b/unianimate/dwpose/onnxdet.py deleted file mode 100644 index 15ae797..0000000 --- a/unianimate/dwpose/onnxdet.py +++ /dev/null @@ -1,127 +0,0 @@ -import cv2 -import numpy as np - -import onnxruntime - -def nms(boxes, scores, nms_thr): - """Single class NMS implemented in Numpy.""" - x1 = boxes[:, 0] - y1 = boxes[:, 1] - x2 = boxes[:, 2] - y2 = boxes[:, 3] - - areas = (x2 - x1 + 1) * (y2 - y1 + 1) - order = scores.argsort()[::-1] - - keep = [] - while order.size > 0: - i = order[0] - keep.append(i) - xx1 = np.maximum(x1[i], x1[order[1:]]) - yy1 = np.maximum(y1[i], y1[order[1:]]) - xx2 = np.minimum(x2[i], x2[order[1:]]) - yy2 = np.minimum(y2[i], y2[order[1:]]) - - w = np.maximum(0.0, xx2 - xx1 + 1) - h = np.maximum(0.0, yy2 - yy1 + 1) - inter = w * h - ovr = inter / (areas[i] + areas[order[1:]] - inter) - - inds = np.where(ovr <= nms_thr)[0] - order = order[inds + 1] - - return keep - -def multiclass_nms(boxes, scores, nms_thr, score_thr): - """Multiclass NMS implemented in Numpy. Class-aware version.""" - final_dets = [] - num_classes = scores.shape[1] - for cls_ind in range(num_classes): - cls_scores = scores[:, cls_ind] - valid_score_mask = cls_scores > score_thr - if valid_score_mask.sum() == 0: - continue - else: - valid_scores = cls_scores[valid_score_mask] - valid_boxes = boxes[valid_score_mask] - keep = nms(valid_boxes, valid_scores, nms_thr) - if len(keep) > 0: - cls_inds = np.ones((len(keep), 1)) * cls_ind - dets = np.concatenate( - [valid_boxes[keep], valid_scores[keep, None], cls_inds], 1 - ) - final_dets.append(dets) - if len(final_dets) == 0: - return None - return np.concatenate(final_dets, 0) - -def demo_postprocess(outputs, img_size, p6=False): - grids = [] - expanded_strides = [] - strides = [8, 16, 32] if not p6 else [8, 16, 32, 64] - - hsizes = [img_size[0] // stride for stride in strides] - wsizes = [img_size[1] // stride for stride in strides] - - for hsize, wsize, stride in zip(hsizes, wsizes, strides): - xv, yv = np.meshgrid(np.arange(wsize), np.arange(hsize)) - grid = np.stack((xv, yv), 2).reshape(1, -1, 2) - grids.append(grid) - shape = grid.shape[:2] - expanded_strides.append(np.full((*shape, 1), stride)) - - grids = np.concatenate(grids, 1) - expanded_strides = np.concatenate(expanded_strides, 1) - outputs[..., :2] = (outputs[..., :2] + grids) * expanded_strides - outputs[..., 2:4] = np.exp(outputs[..., 2:4]) * expanded_strides - - return outputs - -def preprocess(img, input_size, swap=(2, 0, 1)): - if len(img.shape) == 3: - padded_img = np.ones((input_size[0], input_size[1], 3), dtype=np.uint8) * 114 - else: - padded_img = np.ones(input_size, dtype=np.uint8) * 114 - - r = min(input_size[0] / img.shape[0], input_size[1] / img.shape[1]) - resized_img = cv2.resize( - img, - (int(img.shape[1] * r), int(img.shape[0] * r)), - interpolation=cv2.INTER_LINEAR, - ).astype(np.uint8) - padded_img[: int(img.shape[0] * r), : int(img.shape[1] * r)] = resized_img - - padded_img = padded_img.transpose(swap) - padded_img = np.ascontiguousarray(padded_img, dtype=np.float32) - return padded_img, r - -def inference_detector(session, oriImg): - input_shape = (640,640) - img, ratio = preprocess(oriImg, input_shape) - - ort_inputs = {session.get_inputs()[0].name: img[None, :, :, :]} - - output = session.run(None, ort_inputs) - - predictions = demo_postprocess(output[0], input_shape)[0] - - boxes = predictions[:, :4] - scores = predictions[:, 4:5] * predictions[:, 5:] - - boxes_xyxy = np.ones_like(boxes) - boxes_xyxy[:, 0] = boxes[:, 0] - boxes[:, 2]/2. - boxes_xyxy[:, 1] = boxes[:, 1] - boxes[:, 3]/2. - boxes_xyxy[:, 2] = boxes[:, 0] + boxes[:, 2]/2. - boxes_xyxy[:, 3] = boxes[:, 1] + boxes[:, 3]/2. - boxes_xyxy /= ratio - dets = multiclass_nms(boxes_xyxy, scores, nms_thr=0.45, score_thr=0.1) - if dets is not None: - final_boxes, final_scores, final_cls_inds = dets[:, :4], dets[:, 4], dets[:, 5] - isscore = final_scores>0.3 - iscat = final_cls_inds == 0 - isbbox = [ i and j for (i, j) in zip(isscore, iscat)] - final_boxes = final_boxes[isbbox] - else: - final_boxes = np.array([]) - - return final_boxes diff --git a/unianimate/dwpose/onnxpose.py b/unianimate/dwpose/onnxpose.py deleted file mode 100644 index 79cd4a0..0000000 --- a/unianimate/dwpose/onnxpose.py +++ /dev/null @@ -1,360 +0,0 @@ -from typing import List, Tuple - -import cv2 -import numpy as np -import onnxruntime as ort - -def preprocess( - img: np.ndarray, out_bbox, input_size: Tuple[int, int] = (192, 256) -) -> Tuple[np.ndarray, np.ndarray, np.ndarray]: - """Do preprocessing for RTMPose model inference. - - Args: - img (np.ndarray): Input image in shape. - input_size (tuple): Input image size in shape (w, h). - - Returns: - tuple: - - resized_img (np.ndarray): Preprocessed image. - - center (np.ndarray): Center of image. - - scale (np.ndarray): Scale of image. - """ - # get shape of image - img_shape = img.shape[:2] - out_img, out_center, out_scale = [], [], [] - if len(out_bbox) == 0: - out_bbox = [[0, 0, img_shape[1], img_shape[0]]] - for i in range(len(out_bbox)): - x0 = out_bbox[i][0] - y0 = out_bbox[i][1] - x1 = out_bbox[i][2] - y1 = out_bbox[i][3] - bbox = np.array([x0, y0, x1, y1]) - - # get center and scale - center, scale = bbox_xyxy2cs(bbox, padding=1.25) - - # do affine transformation - resized_img, scale = top_down_affine(input_size, scale, center, img) - - # normalize image - mean = np.array([123.675, 116.28, 103.53]) - std = np.array([58.395, 57.12, 57.375]) - resized_img = (resized_img - mean) / std - - out_img.append(resized_img) - out_center.append(center) - out_scale.append(scale) - - return out_img, out_center, out_scale - - -def inference(sess: ort.InferenceSession, img: np.ndarray) -> np.ndarray: - """Inference RTMPose model. - - Args: - sess (ort.InferenceSession): ONNXRuntime session. - img (np.ndarray): Input image in shape. - - Returns: - outputs (np.ndarray): Output of RTMPose model. - """ - all_out = [] - # build input - for i in range(len(img)): - input = [img[i].transpose(2, 0, 1)] - - # build output - sess_input = {sess.get_inputs()[0].name: input} - sess_output = [] - for out in sess.get_outputs(): - sess_output.append(out.name) - - # run model - outputs = sess.run(sess_output, sess_input) - all_out.append(outputs) - - return all_out - - -def postprocess(outputs: List[np.ndarray], - model_input_size: Tuple[int, int], - center: Tuple[int, int], - scale: Tuple[int, int], - simcc_split_ratio: float = 2.0 - ) -> Tuple[np.ndarray, np.ndarray]: - """Postprocess for RTMPose model output. - - Args: - outputs (np.ndarray): Output of RTMPose model. - model_input_size (tuple): RTMPose model Input image size. - center (tuple): Center of bbox in shape (x, y). - scale (tuple): Scale of bbox in shape (w, h). - simcc_split_ratio (float): Split ratio of simcc. - - Returns: - tuple: - - keypoints (np.ndarray): Rescaled keypoints. - - scores (np.ndarray): Model predict scores. - """ - all_key = [] - all_score = [] - for i in range(len(outputs)): - # use simcc to decode - simcc_x, simcc_y = outputs[i] - keypoints, scores = decode(simcc_x, simcc_y, simcc_split_ratio) - - # rescale keypoints - keypoints = keypoints / model_input_size * scale[i] + center[i] - scale[i] / 2 - all_key.append(keypoints[0]) - all_score.append(scores[0]) - - return np.array(all_key), np.array(all_score) - - -def bbox_xyxy2cs(bbox: np.ndarray, - padding: float = 1.) -> Tuple[np.ndarray, np.ndarray]: - """Transform the bbox format from (x,y,w,h) into (center, scale) - - Args: - bbox (ndarray): Bounding box(es) in shape (4,) or (n, 4), formatted - as (left, top, right, bottom) - padding (float): BBox padding factor that will be multilied to scale. - Default: 1.0 - - Returns: - tuple: A tuple containing center and scale. - - np.ndarray[float32]: Center (x, y) of the bbox in shape (2,) or - (n, 2) - - np.ndarray[float32]: Scale (w, h) of the bbox in shape (2,) or - (n, 2) - """ - # convert single bbox from (4, ) to (1, 4) - dim = bbox.ndim - if dim == 1: - bbox = bbox[None, :] - - # get bbox center and scale - x1, y1, x2, y2 = np.hsplit(bbox, [1, 2, 3]) - center = np.hstack([x1 + x2, y1 + y2]) * 0.5 - scale = np.hstack([x2 - x1, y2 - y1]) * padding - - if dim == 1: - center = center[0] - scale = scale[0] - - return center, scale - - -def _fix_aspect_ratio(bbox_scale: np.ndarray, - aspect_ratio: float) -> np.ndarray: - """Extend the scale to match the given aspect ratio. - - Args: - scale (np.ndarray): The image scale (w, h) in shape (2, ) - aspect_ratio (float): The ratio of ``w/h`` - - Returns: - np.ndarray: The reshaped image scale in (2, ) - """ - w, h = np.hsplit(bbox_scale, [1]) - bbox_scale = np.where(w > h * aspect_ratio, - np.hstack([w, w / aspect_ratio]), - np.hstack([h * aspect_ratio, h])) - return bbox_scale - - -def _rotate_point(pt: np.ndarray, angle_rad: float) -> np.ndarray: - """Rotate a point by an angle. - - Args: - pt (np.ndarray): 2D point coordinates (x, y) in shape (2, ) - angle_rad (float): rotation angle in radian - - Returns: - np.ndarray: Rotated point in shape (2, ) - """ - sn, cs = np.sin(angle_rad), np.cos(angle_rad) - rot_mat = np.array([[cs, -sn], [sn, cs]]) - return rot_mat @ pt - - -def _get_3rd_point(a: np.ndarray, b: np.ndarray) -> np.ndarray: - """To calculate the affine matrix, three pairs of points are required. This - function is used to get the 3rd point, given 2D points a & b. - - The 3rd point is defined by rotating vector `a - b` by 90 degrees - anticlockwise, using b as the rotation center. - - Args: - a (np.ndarray): The 1st point (x,y) in shape (2, ) - b (np.ndarray): The 2nd point (x,y) in shape (2, ) - - Returns: - np.ndarray: The 3rd point. - """ - direction = a - b - c = b + np.r_[-direction[1], direction[0]] - return c - - -def get_warp_matrix(center: np.ndarray, - scale: np.ndarray, - rot: float, - output_size: Tuple[int, int], - shift: Tuple[float, float] = (0., 0.), - inv: bool = False) -> np.ndarray: - """Calculate the affine transformation matrix that can warp the bbox area - in the input image to the output size. - - Args: - center (np.ndarray[2, ]): Center of the bounding box (x, y). - scale (np.ndarray[2, ]): Scale of the bounding box - wrt [width, height]. - rot (float): Rotation angle (degree). - output_size (np.ndarray[2, ] | list(2,)): Size of the - destination heatmaps. - shift (0-100%): Shift translation ratio wrt the width/height. - Default (0., 0.). - inv (bool): Option to inverse the affine transform direction. - (inv=False: src->dst or inv=True: dst->src) - - Returns: - np.ndarray: A 2x3 transformation matrix - """ - shift = np.array(shift) - src_w = scale[0] - dst_w = output_size[0] - dst_h = output_size[1] - - # compute transformation matrix - rot_rad = np.deg2rad(rot) - src_dir = _rotate_point(np.array([0., src_w * -0.5]), rot_rad) - dst_dir = np.array([0., dst_w * -0.5]) - - # get four corners of the src rectangle in the original image - src = np.zeros((3, 2), dtype=np.float32) - src[0, :] = center + scale * shift - src[1, :] = center + src_dir + scale * shift - src[2, :] = _get_3rd_point(src[0, :], src[1, :]) - - # get four corners of the dst rectangle in the input image - dst = np.zeros((3, 2), dtype=np.float32) - dst[0, :] = [dst_w * 0.5, dst_h * 0.5] - dst[1, :] = np.array([dst_w * 0.5, dst_h * 0.5]) + dst_dir - dst[2, :] = _get_3rd_point(dst[0, :], dst[1, :]) - - if inv: - warp_mat = cv2.getAffineTransform(np.float32(dst), np.float32(src)) - else: - warp_mat = cv2.getAffineTransform(np.float32(src), np.float32(dst)) - - return warp_mat - - -def top_down_affine(input_size: dict, bbox_scale: dict, bbox_center: dict, - img: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: - """Get the bbox image as the model input by affine transform. - - Args: - input_size (dict): The input size of the model. - bbox_scale (dict): The bbox scale of the img. - bbox_center (dict): The bbox center of the img. - img (np.ndarray): The original image. - - Returns: - tuple: A tuple containing center and scale. - - np.ndarray[float32]: img after affine transform. - - np.ndarray[float32]: bbox scale after affine transform. - """ - w, h = input_size - warp_size = (int(w), int(h)) - - # reshape bbox to fixed aspect ratio - bbox_scale = _fix_aspect_ratio(bbox_scale, aspect_ratio=w / h) - - # get the affine matrix - center = bbox_center - scale = bbox_scale - rot = 0 - warp_mat = get_warp_matrix(center, scale, rot, output_size=(w, h)) - - # do affine transform - img = cv2.warpAffine(img, warp_mat, warp_size, flags=cv2.INTER_LINEAR) - - return img, bbox_scale - - -def get_simcc_maximum(simcc_x: np.ndarray, - simcc_y: np.ndarray) -> Tuple[np.ndarray, np.ndarray]: - """Get maximum response location and value from simcc representations. - - Note: - instance number: N - num_keypoints: K - heatmap height: H - heatmap width: W - - Args: - simcc_x (np.ndarray): x-axis SimCC in shape (K, Wx) or (N, K, Wx) - simcc_y (np.ndarray): y-axis SimCC in shape (K, Wy) or (N, K, Wy) - - Returns: - tuple: - - locs (np.ndarray): locations of maximum heatmap responses in shape - (K, 2) or (N, K, 2) - - vals (np.ndarray): values of maximum heatmap responses in shape - (K,) or (N, K) - """ - N, K, Wx = simcc_x.shape - simcc_x = simcc_x.reshape(N * K, -1) - simcc_y = simcc_y.reshape(N * K, -1) - - # get maximum value locations - x_locs = np.argmax(simcc_x, axis=1) - y_locs = np.argmax(simcc_y, axis=1) - locs = np.stack((x_locs, y_locs), axis=-1).astype(np.float32) - max_val_x = np.amax(simcc_x, axis=1) - max_val_y = np.amax(simcc_y, axis=1) - - # get maximum value across x and y axis - mask = max_val_x > max_val_y - max_val_x[mask] = max_val_y[mask] - vals = max_val_x - locs[vals <= 0.] = -1 - - # reshape - locs = locs.reshape(N, K, 2) - vals = vals.reshape(N, K) - - return locs, vals - - -def decode(simcc_x: np.ndarray, simcc_y: np.ndarray, - simcc_split_ratio) -> Tuple[np.ndarray, np.ndarray]: - """Modulate simcc distribution with Gaussian. - - Args: - simcc_x (np.ndarray[K, Wx]): model predicted simcc in x. - simcc_y (np.ndarray[K, Wy]): model predicted simcc in y. - simcc_split_ratio (int): The split ratio of simcc. - - Returns: - tuple: A tuple containing center and scale. - - np.ndarray[float32]: keypoints in shape (K, 2) or (n, K, 2) - - np.ndarray[float32]: scores in shape (K,) or (n, K) - """ - keypoints, scores = get_simcc_maximum(simcc_x, simcc_y) - keypoints /= simcc_split_ratio - - return keypoints, scores - - -def inference_pose(session, out_bbox, oriImg): - h, w = session.get_inputs()[0].shape[2:] - model_input_size = (w, h) - resized_img, center, scale = preprocess(oriImg, out_bbox, model_input_size) - outputs = inference(session, resized_img) - keypoints, scores = postprocess(outputs, model_input_size, center, scale) - - return keypoints, scores \ No newline at end of file diff --git a/unianimate/dwpose/util.py b/unianimate/dwpose/util.py index f721442..c8aa250 100644 --- a/unianimate/dwpose/util.py +++ b/unianimate/dwpose/util.py @@ -180,12 +180,11 @@ def draw_body_and_foot(canvas, candidate, subset, score, stick_width=4, draw_bod # Append head elements based on the condition limbSeq_and_colors += head_elements - for limb_info in limbSeq_and_colors[:17]: + for limb_info in limbSeq_and_colors[:19]: limbSeq, color = limb_info for n in range(len(subset)): index = subset[n][np.array(limbSeq) - 1] - conf = score[n][np.array(limbSeq) - 1] - if conf[0] < 0.3 or conf[1] < 0.3: + if index[0] < 0.3 or index[1] < 0.3: continue Y = candidate[index.astype(int), 0] * float(W) X = candidate[index.astype(int), 1] * float(H) @@ -194,11 +193,11 @@ def draw_body_and_foot(canvas, candidate, subset, score, stick_width=4, draw_bod length = np.sqrt((X[0] - X[1]) ** 2 + (Y[0] - Y[1]) ** 2) angle = math.degrees(math.atan2(X[0] - X[1], Y[0] - Y[1])) polygon = cv2.ellipse2Poly((int(mY), int(mX)), (int(length / 2), stick_width), int(angle), 0, 360, 1) - cv2.fillConvexPoly(canvas, polygon, alpha_blend_color(color, conf[0] * conf[1])) + cv2.fillConvexPoly(canvas, polygon, alpha_blend_color(color, index[0] * index[1])) canvas = (canvas * 0.6).astype(np.uint8) - for limb_info in limbSeq_and_colors[:18]: + for limb_info in limbSeq_and_colors[:19]: limbSeq, color = limb_info for i in limbSeq: for n in range(len(subset)): @@ -209,12 +208,9 @@ def draw_body_and_foot(canvas, candidate, subset, score, stick_width=4, draw_bod conf = score[n][i - 1] if not np.isfinite(x) or not np.isfinite(y): continue - x = int(np.clip(x * W, 0, W - 1 - )) + x = int(np.clip(x * W, 0, W - 1)) y = int(np.clip(y * H, 0, H - 1)) - # x = int(x * W) - # y = int(y * H) - cv2.circle(canvas, (x, y), 4, alpha_blend_color(color, conf), thickness=-1) + cv2.circle(canvas, (x, y), body_keypoint_size, alpha_blend_color(color, conf), thickness=-1) return canvas diff --git a/unianimate/dwpose/wholebody.py b/unianimate/dwpose/wholebody.py index b838661..ab2022c 100644 --- a/unianimate/dwpose/wholebody.py +++ b/unianimate/dwpose/wholebody.py @@ -1,7 +1,6 @@ import numpy as np from .jit_det import inference_detector as inference_jit_yolox from .jit_pose import inference_pose as inference_jit_pose -import os class Wholebody: diff --git a/unianimate/nodes.py b/unianimate/nodes.py index eb54de4..9fb97d0 100644 --- a/unianimate/nodes.py +++ b/unianimate/nodes.py @@ -1,9 +1,14 @@ - +import torch import torch.nn as nn +import os, copy, math +import numpy as np +from tqdm import tqdm + from ..utils import log + import comfy.model_management as mm from comfy.utils import ProgressBar -from tqdm import tqdm + def update_transformer(transformer, state_dict): @@ -56,14 +61,6 @@ def update_transformer(transformer, state_dict): # 3rd Edited by ControlNet # 4th Edited by ControlNet (added face and correct hands) -import os -import torch -import numpy as np -import copy -import torch -import numpy as np -import math - from .dwpose.wholebody import Wholebody def smoothing_factor(t_e, cutoff): @@ -772,10 +769,14 @@ class WanVideoUniAnimateDWPoseDetector: ref = reference_pose_image ref_np = ref.cpu().numpy() * 255 + prev_fuser_state = torch._C._jit_texpr_fuser_enabled() + torch._C._jit_set_texpr_fuser_enabled(False) # removes warmup delay, may want to enable later poses, reference_pose = pose_extract(pose_np, ref_np, self.dwpose_detector, height, width, score_threshold, stick_width=stick_width, draw_body=draw_body, body_keypoint_size=body_keypoint_size, draw_feet=draw_feet, draw_hands=draw_hands, hand_keypoint_size=hand_keypoint_size, handle_not_detected=handle_not_detected, draw_head=draw_head) poses = poses / 255.0 + torch._C._jit_set_texpr_fuser_enabled(prev_fuser_state) + if reference_pose_image is not None: reference_pose = reference_pose.unsqueeze(0) / 255.0 else: From 9076a3aae8b9193133dd626829cd494ed15267b7 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 23 Aug 2025 20:57:41 +0300 Subject: [PATCH 16/16] bump version --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 1aa8ce0..3567161 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "ComfyUI-WanVideoWrapper" description = "ComfyUI wrapper nodes for WanVideo" -version = "1.2.9" +version = "1.3.0" license = {file = "LICENSE"} dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.33.0", "peft >= 0.15.0", "ftfy", "gguf >= 0.14.0", "pyloudnorm"]