From 04534e5eed1ba305aee14dea050e450bb3cf4b4e Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 4 Apr 2025 15:23:39 +0300 Subject: [PATCH] Allow using multiple VACE embeds, fix VACE blockswap error when vace blocks not set and blockswap is set Experimental --- nodes.py | 46 ++++++++++++++++++++++++--------------- wanvideo/modules/model.py | 35 ++++++++++++++++++++--------- 2 files changed, 54 insertions(+), 27 deletions(-) diff --git a/nodes.py b/nodes.py index 25183a2..5993c87 100644 --- a/nodes.py +++ b/nodes.py @@ -1664,6 +1664,7 @@ class WanVideoVACEEncode: "input_frames": ("IMAGE",), "ref_images": ("IMAGE",), "input_masks": ("MASK",), + "prev_vace_embeds": ("WANVIDIMAGE_EMBEDS",), }, } @@ -1672,7 +1673,7 @@ class WanVideoVACEEncode: FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, vae, width, height, num_frames, strength, vace_start_percent, vace_end_percent, input_frames=None, ref_images=None, input_masks=None): + def process(self, vae, width, height, num_frames, strength, vace_start_percent, vace_end_percent, input_frames=None, ref_images=None, input_masks=None, prev_vace_embeds=None): self.device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -1742,7 +1743,11 @@ class WanVideoVACEEncode: "target_shape": target_shape, "vace_start_percent": vace_start_percent, "vace_end_percent": vace_end_percent, + "additional_vace_inputs": [], } + + if prev_vace_embeds is not None: + vace_input["additional_vace_inputs"].append(prev_vace_embeds) return (vace_input,) def vace_encode_frames(self, frames, ref_images, masks=None): @@ -1867,7 +1872,7 @@ class WanVideoVACEStartToEndFrame: return (out_batch.cpu().float(), masks.cpu().float()) -#region Sampler + class WanVideoContextOptions: @classmethod @@ -1973,6 +1978,7 @@ class WanVideoExperimentalArgs: def process(self, **kwargs): return (kwargs,) +#region Sampler class WanVideoSampler: @classmethod def INPUT_TYPES(s): @@ -2064,7 +2070,7 @@ class WanVideoSampler: seed_g.manual_seed(seed) control_latents, clip_fea, clip_fea_neg, end_image = None, None, None, None - vace_context, vace_scale = None, None + vace_data, vace_scale = None, None fun_model, has_ref = False, False image_cond = image_embeds.get("image_embeds", None) @@ -2108,6 +2114,20 @@ class WanVideoSampler: vace_start_percent = image_embeds.get("vace_start_percent", 0.0) vace_end_percent = image_embeds.get("vace_end_percent", 1.0) + vace_additional_embeds = image_embeds.get("additional_vace_inputs", []) + + vace_data = [ + {"context": vace_context, "scale": vace_scale, "start": vace_start_percent, "end": vace_end_percent} + ] + if len(vace_additional_embeds) > 0: + for i in range(len(vace_additional_embeds)): + vace_data.append({ + "context": vace_additional_embeds[i]["vace_context"], + "scale": vace_additional_embeds[i]["vace_scale"], + "start": vace_additional_embeds[i]["vace_start_percent"], + "end": vace_additional_embeds[i]["vace_end_percent"], + }) + noise = torch.randn( target_shape[0], target_shape[1] + 1 if has_ref else target_shape[1], @@ -2361,7 +2381,7 @@ class WanVideoSampler: zero_star_steps = experimental_args.get("zero_star_steps", 0) #region model pred - def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_context=None, teacache_state=None): + def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_data=None, teacache_state=None): with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init: @@ -2396,13 +2416,6 @@ class WanVideoSampler: patcher.model.is_patched = True else: image_cond_input = image_cond - - if vace_context is not None: - vace_context_input = vace_context - if not (vace_start_percent <= current_step_percentage <= vace_end_percent) or \ - (vace_end_percent < 1.0 and vace_end_percent > 0 and idx == 0 and - (current_step_percentage >= vace_start_percent and current_step_percentage > vace_end_percent)): - vace_context_input = None base_params = { 'seq_len': seq_len, @@ -2412,8 +2425,7 @@ class WanVideoSampler: 'current_step': idx, 'y': [image_cond_input] if image_cond_input is not None else None, 'control_lora_enabled': control_lora_enabled, - 'vace_context': vace_context_input if vace_context is not None else None, - 'vace_scale': vace_scale, + 'vace_data': vace_data if vace_data is not None else None, } batch_size = 1 @@ -2710,10 +2722,10 @@ class WanVideoSampler: partial_img_emb[:, 0, :, :] = partial_image_cond partial_vace_context = None - if vace_context is not None: - partial_vace_context = vace_context[0][:, c, :, :] + if vace_data is not None: + partial_vace_context = vace_data[0]["context"][0][:, c, :, :] if has_ref: - partial_vace_context[:, 0, :, :] = vace_context[0][:, 0, :, :] + partial_vace_context[:, 0, :, :] = vace_data[0]["context"][0][:, 0, :, :] partial_vace_context = [partial_vace_context] partial_latent_model_input = latent_model_input[:, c, :, :] @@ -2744,7 +2756,7 @@ class WanVideoSampler: cfg[idx], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"], - timestep, idx, image_cond, clip_fea, control_latents, vace_context, + timestep, idx, image_cond, clip_fea, control_latents, vace_data, teacache_state=self.teacache_state) if latent_shift_loop: diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index ee32d5a..d2a4f87 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -556,13 +556,14 @@ class BaseWanAttentionBlock(WanAttentionBlock): super().__init__(cross_attn_type, dim, ffn_dim, num_heads, window_size, qk_norm, cross_attn_norm, eps, attention_mode) self.block_id = block_id - def forward(self, x, vace_hints=None, vace_context_scale=1.0, **kwargs): + def forward(self, x, vace_hints=None, vace_context_scale=[1.0], **kwargs): x = super().forward(x, **kwargs) if vace_hints is None: return x if self.block_id is not None: - x = x + vace_hints[self.block_id].to(x.device) * vace_context_scale + for i in range(len(vace_hints)): + x = x + vace_hints[i][self.block_id].to(x.device) * vace_context_scale[i] return x class Head(nn.Module): @@ -817,6 +818,9 @@ class WanModel(ModelMixin, ConfigMixin): block.to(self.offload_device, non_blocking=self.use_non_blocking) total_offload_memory += block_memory + if blocks_to_swap > 0 and vace_blocks_to_swap == 0: + vace_blocks_to_swap = 1 + if vace_blocks_to_swap > 0 and self.vace_layers is not None: self.vace_blocks_to_swap = vace_blocks_to_swap @@ -889,9 +893,7 @@ class WanModel(ModelMixin, ConfigMixin): current_step=0, pred_id=None, control_lora_enabled=False, - vace_context = None, - vace_scale=1.0, - + vace_data = None, ): r""" Forward pass through the diffusion model @@ -1051,11 +1053,24 @@ class WanModel(ModelMixin, ConfigMixin): video_attention_split_steps=self.video_attention_split_steps ) - if vace_context is not None: - vace_hints = self.forward_vace(x, vace_context, seq_len, kwargs) - vace_context_scale = vace_scale - kwargs['vace_hints'] = vace_hints - kwargs['vace_context_scale'] = vace_context_scale + if vace_data is not None: + vace_hint_list = [] + vace_scale_list = [] + if isinstance(vace_data[0], dict): + for data in vace_data: + if (data["start"] <= current_step_percentage <= data["end"]) or \ + (data["end"] > 0 and current_step == 0 and current_step_percentage >= data["start"]): + + vace_hints = self.forward_vace(x, data["context"], seq_len, kwargs) + vace_hint_list.append(vace_hints) + vace_scale_list.append(data["scale"]) + else: + vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs) + vace_hint_list.append(vace_hints) + vace_scale_list.append(1.0) + + kwargs['vace_hints'] = vace_hint_list + kwargs['vace_context_scale'] = vace_scale_list for b, block in enumerate(self.blocks): if self.slg_blocks is not None: