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()