Support Wan22 VAE in WanVideoLatentReScale

This commit is contained in:
kijai
2025-08-22 15:04:29 +03:00
parent bb5503c9bb
commit b1ceaaeb2e
3 changed files with 113 additions and 20 deletions
+51 -17
View File
@@ -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)
+1 -1
View File
@@ -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):
+61 -2
View File
@@ -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"
}