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