diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 9b32b99..af9ea79 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1046,6 +1046,29 @@ class WanVideoSetAttentionModeOverride: return (model_clone,) +class WanVideoUltraVicoSettings: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("WANVIDEOMODEL", ), + "alpha": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.001}), + }, + } + + RETURN_TYPES = ("WANVIDEOMODEL",) + RETURN_NAMES = ("model", ) + FUNCTION = "getmodelpath" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Override the attention mode for the model for specific step and/or block range" + + def getmodelpath(self, model, alpha): + model_clone = model.clone() + model_clone.model_options['transformer_options']["ultravico_alpha"] = alpha + + return (model_clone,) + + #region Model loading class WanVideoModelLoader: @classmethod @@ -2074,6 +2097,7 @@ NODE_CLASS_MAPPINGS = { "LoadWanVideoT5TextEncoder": LoadWanVideoT5TextEncoder, "LoadWanVideoClipTextEncoder": LoadWanVideoClipTextEncoder, "WanVideoSetAttentionModeOverride": WanVideoSetAttentionModeOverride, + "WanVideoUltraVicoSettings": WanVideoUltraVicoSettings, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -2093,4 +2117,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LoadWanVideoT5TextEncoder": "WanVideo T5 Text Encoder Loader", "LoadWanVideoClipTextEncoder": "WanVideo CLIP Text Encoder Loader", "WanVideoSetAttentionModeOverride": "WanVideo Set Attention Mode Override", + "WanVideoUltraVicoSettings": "WanVideo UltraVico Settings" } diff --git a/nodes_sampler.py b/nodes_sampler.py index 6072608..489e216 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -96,7 +96,7 @@ class WanVideoSampler: vae = image_embeds.get("vae", None) tiled_vae = image_embeds.get("tiled_vae", False) - transformer_options = patcher.model_options.get("transformer_options", None) + transformer_options = copy.deepcopy(patcher.model_options.get("transformer_options", None)) merge_loras = transformer_options["merge_loras"] block_swap_args = transformer_options.get("block_swap_args", None) @@ -1414,7 +1414,6 @@ class WanVideoSampler: 'is_uncond': False, # is unconditional 'current_step': idx, # current step 'current_step_percentage': current_step_percentage, # current step percentage - 'attention_mode_override': transformer_options.get("attention_mode_override", None), 'last_step': len(timesteps) - 1 == idx, # is last step 'control_lora_enabled': control_lora_enabled, # control lora toggle for patch embed selection 'enhance_enabled': enhance_enabled, # enhance-a-video toggle @@ -1466,6 +1465,7 @@ class WanVideoSampler: "one_to_all_input": one_to_all_data, # One-to-All input "one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0, "scail_input": scail_data_in, # SCAIL input + "transformer_options": transformer_options } batch_size = 1 diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 3b891d3..f31f6be 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -94,7 +94,7 @@ except: def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0., softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16, - attention_mode='sdpa', attn_mask=None, multi_factor=0.9, frame_tokens=1536, heads=128): + attention_mode='sdpa', attn_mask=None, transformer_options={}, frame_tokens=1536, heads=128): if "flash" in attention_mode: return flash_attention(q, k, v, q_lens=q_lens, k_lens=k_lens, dropout_p=dropout_p, softmax_scale=softmax_scale, q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3, @@ -108,7 +108,7 @@ def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k elif attention_mode == 'sageattn': return sageattn_func(q, k, v, tensor_layout="NHD").contiguous() elif attention_mode == 'sageattn_ultravico': - return sageattn_func_ultravico([q, k, v], multi_factor=multi_factor, frame_tokens=frame_tokens).contiguous() + return sageattn_func_ultravico([q, k, v], multi_factor=transformer_options.get("ultravico_alpha", 0.9), frame_tokens=frame_tokens).contiguous() elif attention_mode == 'comfy': return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True) else: # sdpa diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 4a85818..27d5384 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -467,7 +467,7 @@ class WanSelfAttention(nn.Module): v = (self.v(x) + self.v_loras(x)).view(b, s, n, d) return q, k, v - def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None, onetoall_ref=None, onetoall_ref_scale=1.0, frame_tokens=1536): + def forward(self, q, k, v, seq_lens, transformer_options={}, attention_mode_override=None, lynx_ref_feature=None, lynx_ref_scale=1.0, onetoall_ref=None, onetoall_ref_scale=1.0, frame_tokens=1536): r""" Args: x(Tensor): Shape [B, L, num_heads, C / num_heads] @@ -482,7 +482,7 @@ class WanSelfAttention(nn.Module): if self.ref_adapter is not None and lynx_ref_feature is not None: ref_x = self.ref_adapter(self, q, lynx_ref_feature) - x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads, frame_tokens=frame_tokens) + x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads, frame_tokens=frame_tokens, transformer_options=transformer_options) if self.ref_adapter is not None and lynx_ref_feature is not None: x = x.add(ref_x, alpha=lynx_ref_scale) @@ -1006,7 +1006,7 @@ class WanAttentionBlock(nn.Module): longcat_num_cond_latents=0, longcat_avatar_options=None, #longcat image cond amount x_onetoall_ref=None, onetoall_freqs=None, onetoall_ref=None, onetoall_ref_scale=1.0, #one-to-all e_tr=None, tr_num=0, tr_start=0, #token replacement - attention_mode_override=None, frame_tokens=None, + attention_mode_override=None, frame_tokens=None, transformer_options={} ): r""" Args: @@ -1189,13 +1189,13 @@ class WanAttentionBlock(nn.Module): if longcat_num_cond_latents == 1: num_cond_latents_thw = longcat_num_cond_latents * (N // num_latent_frames) # process the noise tokens - x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens, attention_mode_override=attention_mode_override) + x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # process the condition tokens x_cond = self.self_attn.forward( q[:, :num_cond_latents_thw].contiguous(), k[:, :num_cond_latents_thw].contiguous(), v[:, :num_cond_latents_thw].contiguous(), - seq_lens, attention_mode_override=attention_mode_override) + seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # merge x_cond and x_noise y = torch.cat([x_cond, x_noise], dim=1).contiguous() elif longcat_num_cond_latents > 1: # video continuation @@ -1224,12 +1224,12 @@ class WanAttentionBlock(nn.Module): k_non_ref = k[:, num_ref_latents_thw:].contiguous() v_non_ref = v[:, num_ref_latents_thw:].contiguous() - x_noise_front = self.self_attn.forward(q_noise_front, k, v, seq_lens) # q_front has attention with ref + cond + noisy - x_noise_back = self.self_attn.forward(q_noise_back, k, v, seq_lens) # q_back has attention with ref + cond + noisy - x_noise_maskref = self.self_attn.forward(q_noise_maskref, k_non_ref, v_non_ref, seq_lens) # q_mask has attention with cond+noisy + x_noise_front = self.self_attn.forward(q_noise_front, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_front has attention with ref + cond + noisy + x_noise_back = self.self_attn.forward(q_noise_back, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_back has attention with ref + cond + noisy + x_noise_maskref = self.self_attn.forward(q_noise_maskref, k_non_ref, v_non_ref, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # q_mask has attention with cond+noisy x_noise = torch.cat([x_noise_front, x_noise_maskref, x_noise_back], dim=1).contiguous() else: - x_noise = self.self_attn.forward(q_noise, k, v, seq_lens) + x_noise = self.self_attn.forward(q_noise, k, v, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # process the condition tokens q_ref = q[:, :num_ref_latents_thw].contiguous() k_ref = k[:, :num_ref_latents_thw].contiguous() @@ -1237,14 +1237,14 @@ class WanAttentionBlock(nn.Module): q_cond = q[:, num_ref_latents_thw:num_cond_latents_thw].contiguous() k_cond = k[:, num_ref_latents_thw:num_cond_latents_thw].contiguous() v_cond = v[:, num_ref_latents_thw:num_cond_latents_thw].contiguous() - x_ref = self.self_attn.forward(q_ref, k_ref, v_ref, seq_lens, attention_mode_override=attention_mode_override) - x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens, attention_mode_override=attention_mode_override) + x_ref = self.self_attn.forward(q_ref, k_ref, v_ref, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) + x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens, attention_mode_override=attention_mode_override, transformer_options=transformer_options) # merge x_cond and x_noise y = torch.cat([x_ref, x_cond, x_noise], dim=1).contiguous() else: y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale, - onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale, attention_mode_override=attention_mode_override, frame_tokens=frame_tokens) + onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale, attention_mode_override=attention_mode_override, transformer_options=transformer_options, frame_tokens=frame_tokens) del q, k, v @@ -2281,7 +2281,6 @@ class WanModel(torch.nn.Module): self, x, t, context, seq_len, is_uncond=False, current_step_percentage=0.0, current_step=0, last_step=0, total_steps=50, - attention_mode_override=None, clip_fea=None, y=None, device=torch.device('cuda'), freqs=None, @@ -2320,6 +2319,7 @@ class WanModel(torch.nn.Module): sdancer_input=None, # SteadyDancer one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All scail_input=None, # SCAIL pose + transformer_options={}, ): r""" Forward pass through the diffusion model @@ -3069,6 +3069,7 @@ class WanModel(torch.nn.Module): e_tr=e0_token_replace if use_token_replace else None, tr_start=token_replace_start, tr_num=replace_token_num, + transformer_options=transformer_options ) if self.audio_model is not None: kwargs['e_ovi'] = e0_ovi.to(self.base_dtype) @@ -3130,6 +3131,7 @@ class WanModel(torch.nn.Module): attn_override_blocks = attention_mode = None attention_mode_override_active = False + attention_mode_override = transformer_options.get("attention_mode_override", None) if attention_mode_override is not None: attn_override_blocks = attention_mode_override.get("blocks", range(len(self.blocks))) if attention_mode_override["start_step"] <= current_step < attention_mode_override["end_step"]: