Allow using multiple VACE embeds, fix VACE blockswap error when vace blocks not set and blockswap is set

Experimental
This commit is contained in:
kijai
2025-04-04 15:23:39 +03:00
parent 78559ef1f4
commit 04534e5eed
2 changed files with 54 additions and 27 deletions
+25 -10
View File
@@ -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: