Support Wan22 VAE in WanVideoLatentReScale
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
@@ -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"
|
||||
}
|
||||
Reference in New Issue
Block a user