diff --git a/nodes.py b/nodes.py index 4ac9b0b..5dbd1c3 100644 --- a/nodes.py +++ b/nodes.py @@ -3214,8 +3214,6 @@ class WanVideoDecode: if drop_last: latents = latents[:, :, :-1] - #if is_looped: - # latents = torch.cat([latents[:, :, :warmup_latent_count],latents], dim=2) if type(vae).__name__ == "TAEHV": images = vae.decode_video(latents.permute(0, 2, 1, 3, 4))[0].permute(1, 0, 2, 3) images = torch.clamp(images, 0.0, 1.0) @@ -3227,34 +3225,30 @@ class WanVideoDecode: images = vae.decode(latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//8, tile_y//8), tile_stride=(tile_stride_x//8, tile_stride_y//8))[0] vae.model.clear_cache() - images = images.cpu() + images = images.cpu().float() if normalization == "minmax": - images = (images - images.min()) / (images.max() - images.min()) + images.sub_(images.min()).div_(images.max() - images.min()) else: - images = torch.clamp(images, -1.0, 1.0) - images = (images + 1.0) / 2.0 + images.clamp_(-1.0, 1.0) + images.add_(1.0).div_(2.0) if is_looped: - #images = images[:, warmup_latent_count * 4:] temp_latents = torch.cat([latents[:, :, -3:]] + [latents[:, :, :2]], dim=2) temp_images = vae.decode(temp_latents, device=device, end_=(end_image is not None), tiled=enable_vae_tiling, tile_size=(tile_x//vae.upsampling_factor, tile_y//vae.upsampling_factor), tile_stride=(tile_stride_x//vae.upsampling_factor, tile_stride_y//vae.upsampling_factor))[0] temp_images = (temp_images - temp_images.min()) / (temp_images.max() - temp_images.min()) images = torch.cat([temp_images[:, 9:].to(images), images[:, 5:]], dim=1) if end_image is not None: - #end_image = (end_image - end_image.min()) / (end_image.max() - end_image.min()) - #image[:, -1] = end_image[:, 0].to(image) #not sure about this images = images[:, 0:-1] vae.model.clear_cache() vae.to(offload_device) mm.soft_empty_cache() - images = torch.clamp(images, 0.0, 1.0) - images = images.permute(1, 2, 3, 0).float() + images.clamp_(0.0, 1.0) - return (images,) + return (images.permute(1, 2, 3, 0),) #region VideoEncode class WanVideoEncode: @@ -3295,35 +3289,7 @@ class WanVideoEncode: if image.shape[-1] == 4: image = image[..., :3] - image = image.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W - - - # empty_frame_indices = [] - # for i in range(image.shape[2]): - # if is_image_black(image[:, :, i]): - # empty_frame_indices.append(i) - # empty_frame_indices = [] - # for i in range(image.shape[2]): - # if is_image_black(image[:, :, i]): - # empty_frame_indices.append(i) - # empty_latent_indices = [] - # if empty_frame_indices: - # frames_per_latent = 4 - # num_frames = image.shape[2] - # # Special mapping: latent 0 = [0], latent 1 = [1,2,3,4], latent 2 = [5,6,7,8], ... - # latent_frame_ranges = [] - # latent_frame_ranges.append([0]) - # for i in range(1, math.ceil((num_frames - 1) / frames_per_latent) + 1): - # start = 1 + (i - 1) * frames_per_latent - # end = min(start + frames_per_latent, num_frames) - # latent_frame_ranges.append(list(range(start, end))) - # for latent_idx, latent_frames in enumerate(latent_frame_ranges): - # print(f"latent {latent_idx}: frames {latent_frames}") - # if latent_frames and set(latent_frames).issubset(empty_frame_indices): - # empty_latent_indices.append(latent_idx) - # if empty_latent_indices: - # log.info(f"Empty frames {empty_frame_indices} map to latents {empty_latent_indices}") - + image = image.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W if noise_aug_strength > 0.0: image = add_noise_to_reference_video(image, ratio=noise_aug_strength) @@ -3342,9 +3308,6 @@ class WanVideoEncode: if mask is None: vae.to(offload_device) else: - #latent_mask = mask.clone().to(vae.dtype).to(device) * 2.0 - 1.0 - #latent_mask = latent_mask.unsqueeze(0).unsqueeze(0).repeat(1, 3, 1, 1, 1) - #latent_mask = vae.encode(latent_mask, device=device, tiled=enable_vae_tiling, tile_size=(tile_x, tile_y), tile_stride=(tile_stride_x, tile_stride_y)) target_h, target_w = latents.shape[3:] mask = torch.nn.functional.interpolate( @@ -3362,6 +3325,45 @@ class WanVideoEncode: return ({"samples": latents, "mask": latent_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"] + + 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, @@ -3390,6 +3392,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoBlockList": WanVideoBlockList, "WanVideoTextEncodeCached": WanVideoTextEncodeCached, "WanVideoAddExtraLatent": WanVideoAddExtraLatent, + "WanVideoLatentReScale": WanVideoLatentReScale, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3420,4 +3423,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoBlockList": "WanVideo Block List", "WanVideoTextEncodeCached": "WanVideo TextEncode Cached", "WanVideoAddExtraLatent": "WanVideo Add Extra Latent", + "WanVideoLatentReScale": "WanVideo Latent ReScale", } diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 9f7052c..d6d2e88 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -1,18 +1,21 @@ # Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved. import torch +from ...utils import log +# Flash Attention imports try: import flash_attn_interface FLASH_ATTN_3_AVAILABLE = True -except ModuleNotFoundError: +except Exception as e: FLASH_ATTN_3_AVAILABLE = False try: import flash_attn FLASH_ATTN_2_AVAILABLE = True -except ModuleNotFoundError: +except Exception as e: FLASH_ATTN_2_AVAILABLE = False - + +# Sage Attention imports try: from sageattention import sageattn @torch.compiler.disable() @@ -22,11 +25,11 @@ try: else: return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout=tensor_layout) except Exception as e: - print(f"Warning: Could not load sageattention: {str(e)}") + log.warning(f"Warning: Could not load sageattention: {str(e)}") if isinstance(e, ModuleNotFoundError): - print("sageattention package is not installed") + log.warning("sageattention package is not installed, sageattention will not be available") elif isinstance(e, ImportError) and "DLL" in str(e): - print("sageattention DLL loading error") + log.warning("sageattention DLL loading error, sageattention will not be available") sageattn_func = None try: @@ -34,7 +37,6 @@ try: except: SAGE3_AVAILABLE = False -import warnings __all__ = [ 'flash_attention', @@ -107,9 +109,7 @@ def flash_attention( q = q * q_scale if version is not None and version == 3 and not FLASH_ATTN_3_AVAILABLE: - warnings.warn( - 'Flash attention 3 is not available, use flash attention 2 instead.' - ) + log.warning('Flash attention 3 is not available, use flash attention 2 instead.') # apply attention if (version is None or version == 3) and FLASH_ATTN_3_AVAILABLE: diff --git a/wanvideo/wan_video_vae.py b/wanvideo/wan_video_vae.py index 7a90a3b..881d841 100644 --- a/wanvideo/wan_video_vae.py +++ b/wanvideo/wan_video_vae.py @@ -958,7 +958,9 @@ class VideoVAE_(nn.Module): num_res_blocks=2, attn_scales=[], temperal_downsample=[False, True, True], - dropout=0.0,): + dropout=0.0, + mean=None, + inv_std=None): super().__init__() self.dim = dim self.z_dim = z_dim @@ -967,6 +969,8 @@ class VideoVAE_(nn.Module): self.attn_scales = attn_scales self.temperal_downsample = temperal_downsample self.temperal_upsample = temperal_downsample[::-1] + self.mean = mean + self.inv_std = inv_std # modules self.encoder = Encoder3d(dim, z_dim * 2, dim_mult, num_res_blocks, @@ -984,7 +988,7 @@ class VideoVAE_(nn.Module): #modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames - def encode_2(self, x, scale): + def encode_2(self, x): self.clear_cache() ## cache t = x.shape[2] @@ -1008,18 +1012,14 @@ class VideoVAE_(nn.Module): out = torch.cat([out, out_], 2) out_head = out[:, :, :iter_ - 1, :, :] out_tail = out[:, :, -1, :, :].unsqueeze(2) - mu, log_var = torch.cat([self.conv1(out_head), self.conv1(out_tail)], dim=2).chunk(2, dim=1) - if isinstance(scale[0], torch.Tensor): - scale = [s.to(dtype=mu.dtype, device=mu.device) for s in scale] - mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view( - 1, self.z_dim, 1, 1, 1) - else: - scale = scale.to(dtype=mu.dtype, device=mu.device) - mu = (mu - scale[0]) * scale[1] + mu = torch.cat([self.conv1(out_head), self.conv1(out_tail)], dim=2).chunk(2, dim=1)[0] + + mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu) + return mu - def encode(self, x, scale): + def encode(self, x): self.clear_cache() ## cache t = x.shape[2] @@ -1036,28 +1036,20 @@ class VideoVAE_(nn.Module): feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) out = torch.cat([out, out_], 2) - mu, log_var = self.conv1(out).chunk(2, dim=1) - if isinstance(scale[0], torch.Tensor): - scale = [s.to(dtype=mu.dtype, device=mu.device) for s in scale] - mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view( - 1, self.z_dim, 1, 1, 1) - else: - scale = scale.to(dtype=mu.dtype, device=mu.device) - mu = (mu - scale[0]) * scale[1] + mu = self.conv1(out).chunk(2, dim=1)[0] + + mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu) + return mu #modification originally by @raindrop313 https://github.com/raindrop313/ComfyUI-WanVideoStartEndFrames - def decode_2(self, z, scale): + def decode_2(self, z): self.clear_cache() # z: [b,c,t,h,w] - if isinstance(scale[0], torch.Tensor): - scale = [s.to(dtype=z.dtype, device=z.device) for s in scale] - z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view( - 1, self.z_dim, 1, 1, 1) - else: - scale = scale.to(dtype=z.dtype, device=z.device) - z = z / scale[1] + scale[0] + + z = z / self.inv_std.to(z) + self.mean.to(z) + iter_ = z.shape[2] z_head=z[:,:,:-1,:,:] z_tail=z[:,:,-1,:,:].unsqueeze(2) @@ -1082,17 +1074,13 @@ class VideoVAE_(nn.Module): - def decode(self, z, scale): + def decode(self, z): self.clear_cache() # z: [b,c,t,h,w] pbar = ProgressBar(z.shape[2]) - if isinstance(scale[0], torch.Tensor): - scale = [s.to(dtype=z.dtype, device=z.device) for s in scale] - z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view( - 1, self.z_dim, 1, 1, 1) - else: - scale = scale.to(dtype=z.dtype, device=z.device) - z = z / scale[1] + scale[0] + + z = z / self.inv_std.to(z) + self.mean.to(z) + iter_ = z.shape[2] x = self.conv2(z) for i in range(iter_): @@ -1146,13 +1134,12 @@ class WanVideoVAE(nn.Module): 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 ] - self.mean = torch.tensor(mean) - self.std = torch.tensor(std) - self.scale = [self.mean, 1.0 / self.std] + self.mean = torch.tensor(mean).view(1, z_dim, 1, 1, 1) + self.inv_std = (1.0 / torch.tensor(std)).view(1, z_dim, 1, 1, 1) self.z_dim = z_dim # init model - self.model = VideoVAE_(z_dim=z_dim).eval().requires_grad_(False) + self.model = VideoVAE_(z_dim=z_dim, mean=self.mean, inv_std=self.inv_std).eval().requires_grad_(False) self.upsampling_factor = 8 @@ -1202,7 +1189,7 @@ class WanVideoVAE(nn.Module): pbar = ProgressBar(len(tasks)) for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"): hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device) - hidden_states_batch = self.model.decode(hidden_states_batch, self.scale).to(data_device) + hidden_states_batch = self.model.decode(hidden_states_batch).to(data_device) mask = self.build_mask( hidden_states_batch, @@ -1264,9 +1251,9 @@ class WanVideoVAE(nn.Module): for h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"): hidden_states_batch = video[:, :, :, h:h_, w:w_].to(computation_device) if end_: - hidden_states_batch = self.model.encode_2(hidden_states_batch, self.scale).to(data_device) + hidden_states_batch = self.model.encode_2(hidden_states_batch).to(data_device) else: - hidden_states_batch = self.model.encode(hidden_states_batch, self.scale).to(data_device) + hidden_states_batch = self.model.encode(hidden_states_batch).to(data_device) mask = self.build_mask( hidden_states_batch, @@ -1298,26 +1285,26 @@ class WanVideoVAE(nn.Module): def single_encode(self, video, device): video = video.to(device) - x = self.model.encode(video, self.scale) + x = self.model.encode(video) return x.float() def single_decode(self, hidden_state, device): hidden_state = hidden_state.to(device) - video = self.model.decode(hidden_state, self.scale) - return video.float().clamp_(-1, 1) + video = self.model.decode(hidden_state) + return video def double_encode(self, video, device): print('double_encode') video = video.to(device) - x = self.model.encode_2(video, self.scale) + x = self.model.encode_2(video) return x.float() def double_decode(self, hidden_state, device): print('double_decode') hidden_state = hidden_state.to(device) - video = self.model.decode_2(hidden_state, self.scale) - return video.float().clamp_(-1, 1) + video = self.model.decode_2(hidden_state) + return video def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None): videos = [video.to("cpu") for video in videos] @@ -1384,7 +1371,9 @@ class VideoVAE38_(VideoVAE_): attn_scales=[], temperal_downsample=[False, True, True], dropout=0.0, - dtype=torch.bfloat16): + dtype=torch.bfloat16, + mean=None, + inv_std=None): super(VideoVAE_, self).__init__() self.dim = dim self.z_dim = z_dim @@ -1394,6 +1383,8 @@ class VideoVAE38_(VideoVAE_): self.temperal_downsample = temperal_downsample self.temperal_upsample = temperal_downsample[::-1] self.dtype = dtype + self.mean = mean + self.inv_std = inv_std # modules self.encoder = Encoder3d_38(dim, z_dim * 2, dim_mult, num_res_blocks, @@ -1404,7 +1395,7 @@ class VideoVAE38_(VideoVAE_): attn_scales, self.temperal_upsample, dropout) - def encode(self, x, scale): + def encode(self, x): self.clear_cache() x = patchify(x, patch_size=2) t = x.shape[2] @@ -1420,27 +1411,19 @@ class VideoVAE38_(VideoVAE_): feat_cache=self._enc_feat_map, feat_idx=self._enc_conv_idx) out = torch.cat([out, out_], 2) - mu, log_var = self.conv1(out).chunk(2, dim=1) - if isinstance(scale[0], torch.Tensor): - scale = [s.to(dtype=mu.dtype, device=mu.device) for s in scale] - mu = (mu - scale[0].view(1, self.z_dim, 1, 1, 1)) * scale[1].view( - 1, self.z_dim, 1, 1, 1) - else: - scale = scale.to(dtype=mu.dtype, device=mu.device) - mu = (mu - scale[0]) * scale[1] + mu = self.conv1(out).chunk(2, dim=1)[0] + + mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu) + self.clear_cache() return mu - def decode(self, z, scale): + def decode(self, z): self.clear_cache() - if isinstance(scale[0], torch.Tensor): - scale = [s.to(dtype=z.dtype, device=z.device) for s in scale] - z = z / scale[1].view(1, self.z_dim, 1, 1, 1) + scale[0].view( - 1, self.z_dim, 1, 1, 1) - else: - scale = scale.to(dtype=z.dtype, device=z.device) - z = z / scale[1] + scale[0] + + z = z / self.inv_std.to(z) + self.mean.to(z) + iter_ = z.shape[2] x = self.conv2(z) for i in range(iter_): @@ -1481,12 +1464,11 @@ class WanVideoVAE38(WanVideoVAE): 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 ] - self.mean = torch.tensor(mean) - self.std = torch.tensor(std) - self.scale = [self.mean, 1.0 / self.std] + self.mean = torch.tensor(mean).view(1, z_dim, 1, 1, 1) + self.inv_std = (1.0 / torch.tensor(std)).view(1, z_dim, 1, 1, 1) self.dtype = dtype self.z_dim = z_dim # init model - self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype).eval().requires_grad_(False) + self.model = VideoVAE38_(z_dim=z_dim, dim=dim, dtype=dtype, mean=self.mean, inv_std=self.inv_std).eval().requires_grad_(False) self.upsampling_factor = 16 \ No newline at end of file