Add node for UltraVico params
This commit is contained in:
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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"]:
|
||||
|
||||
Reference in New Issue
Block a user