From 06863fe2d283e89ba2386260ad2700557dda1897 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 12 Aug 2025 11:47:33 +0300 Subject: [PATCH] Allow doing T2V with the I2V encode node when no images are connected and allow using native set latent noise mask -node Latent masking still doesn't work properly --- nodes.py | 70 +++++++++++++++++++++++++-------------- wanvideo/modules/model.py | 27 +++++++++++++++ 2 files changed, 73 insertions(+), 24 deletions(-) diff --git a/nodes.py b/nodes.py index 834d335..fdbb7d3 100644 --- a/nodes.py +++ b/nodes.py @@ -752,7 +752,6 @@ class WanVideoImageToVideoEncode: @classmethod def INPUT_TYPES(s): return {"required": { - "vae": ("WANVAE",), "width": ("INT", {"default": 832, "min": 64, "max": 8096, "step": 8, "tooltip": "Width of the image to encode"}), "height": ("INT", {"default": 480, "min": 64, "max": 8096, "step": 8, "tooltip": "Height of the image to encode"}), "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), @@ -762,6 +761,7 @@ class WanVideoImageToVideoEncode: "force_offload": ("BOOLEAN", {"default": True}), }, "optional": { + "vae": ("WANVAE",), "clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}), "start_image": ("IMAGE", {"tooltip": "Image to encode"}), "end_image": ("IMAGE", {"tooltip": "end frame"}), @@ -779,10 +779,16 @@ class WanVideoImageToVideoEncode: FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, vae, width, height, num_frames, force_offload, noise_aug_strength, + def process(self, width, height, num_frames, force_offload, noise_aug_strength, start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, - temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None): - + temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None): + + if start_image is None and end_image is None: + return WanVideoEmptyEmbeds().process( + num_frames, width, height, control_embeds=control_embeds, extra_latents=extra_latents, + ) + if vae is None: + raise ValueError("VAE is required for image encoding.") H = height W = width @@ -1443,6 +1449,23 @@ class WanVideoFreeInitArgs: def process(self, **kwargs): return (kwargs,) +class WanVideoScheduler: #WIP + @classmethod + def INPUT_TYPES(s): + return {"required": { + "scheduler": (scheduler_list, {"default": "uni_pc"}), + }, + } + + RETURN_TYPES = (scheduler_list, ) + RETURN_NAMES = ("scheduler",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + EXPERIMENTAL = True + + def process(self, scheduler): + return (scheduler,) + #region Sampler class WanVideoSampler: @classmethod @@ -1918,9 +1941,24 @@ class WanVideoSampler: noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples else: noise = input_samples - mask = samples.get("mask", None) + + mask = samples.get("noise_mask", None) if mask is not None: + log.info(f"Latent mask shape: {mask.shape}") original_image = input_samples.to(device) + if len(mask.shape) == 4: + mask = mask.squeeze(1) + + mask = torch.nn.functional.interpolate( + mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W] + size=(noise.shape[1], noise.shape[2], noise.shape[3]), + mode='trilinear', + align_corners=False + ).squeeze(0) # Remove batch dim, keep channel dim + + # Add batch & channel dims for final output + mask = mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1) + if mask.shape[2] != noise.shape[1]: mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2) @@ -3342,27 +3380,10 @@ class WanVideoEncode: if latent_strength != 1.0: latents *= latent_strength - log.info(f"encoded latents shape {latents.shape}") - latent_mask = None - if mask is None: - vae.to(offload_device) - else: - target_h, target_w = latents.shape[3:] - - mask = torch.nn.functional.interpolate( - mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W] - size=(latents.shape[2], target_h, target_w), - mode='trilinear', - align_corners=False - ).squeeze(0) # Remove batch dim, keep channel dim - - # Add batch & channel dims for final output - latent_mask = mask.unsqueeze(0).repeat(1, latents.shape[1], 1, 1, 1) - log.info(f"latent mask shape {latent_mask.shape}") - vae.to(offload_device) + log.info(f"WanVideoEncode: Encoded latents shape {latents.shape}") mm.soft_empty_cache() - return ({"samples": latents, "mask": latent_mask},) + return ({"samples": latents, "noise_mask": mask},) class WanVideoLatentReScale: @classmethod @@ -3432,6 +3453,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoTextEncodeCached": WanVideoTextEncodeCached, "WanVideoAddExtraLatent": WanVideoAddExtraLatent, "WanVideoLatentReScale": WanVideoLatentReScale, + "WanVideoScheduler": WanVideoScheduler } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index f4d77df..f221520 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -448,6 +448,33 @@ class WanSelfAttention(nn.Module): return nag_guidance * nag_alpha + x_positive * (1 - nag_alpha) +class LoRALinearLayer(nn.Module): + def __init__( + self, + in_features: int, + out_features: int, + rank: int = 128, + device=torch.device("cuda"), + dtype=torch.float32, + ): + super().__init__() + self.down = nn.Linear(in_features, rank, bias=False, device=device, dtype=dtype) + self.up = nn.Linear(rank, out_features, bias=False, device=device, dtype=dtype) + self.rank = rank + self.out_features = out_features + self.in_features = in_features + + nn.init.normal_(self.down.weight, std=1 / rank) + nn.init.zeros_(self.up.weight) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + orig_dtype = hidden_states.dtype + dtype = self.down.weight.dtype + + down_hidden_states = self.down(hidden_states.to(dtype)) + up_hidden_states = self.up(down_hidden_states) + return up_hidden_states.to(orig_dtype) + #region crossattn class WanT2VCrossAttention(WanSelfAttention):