Add node for UltraVico params

This commit is contained in:
kijai
2025-12-30 02:53:38 +02:00
parent be41f67fae
commit 3a7100bc39
4 changed files with 44 additions and 17 deletions
+25
View File
@@ -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"
}
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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
+15 -13
View File
@@ -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"]: