Merge branch 'pr/643'

This commit is contained in:
kijai
2025-06-12 18:32:20 +03:00
2 changed files with 195 additions and 63 deletions
+96 -6
View File
@@ -1277,13 +1277,92 @@ class WanVideoTextEncode:
return cleaned_prompt, weights
class WanVideoTextEncodeSingle:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"t5": ("WANTEXTENCODER",),
"prompt": ("STRING", {"default": "", "multiline": True} ),
},
"optional": {
"force_offload": ("BOOLEAN", {"default": True}),
"model_to_offload": ("WANVIDEOMODEL", {"tooltip": "Model to move to offload_device before encoding"}),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
RETURN_NAMES = ("text_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Encodes text prompt into text embedding."
def process(self, t5, prompt, force_offload=True, model_to_offload=None):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
if model_to_offload is not None:
log.info(f"Moving video model to {offload_device}")
model_to_offload.model.to(offload_device)
mm.soft_empty_cache()
encoder = t5["model"]
dtype = t5["dtype"]
encoder.model.to(device)
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=True):
encoded = encoder([prompt], device)
if force_offload:
encoder.model.to(offload_device)
mm.soft_empty_cache()
prompt_embeds_dict = {
"prompt_embeds": encoded,
"negative_prompt_embeds": None,
}
return (prompt_embeds_dict,)
class WanVideoApplyNAG:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"original_text_embeds": ("WANVIDEOTEXTEMBEDS",),
"nag_text_embeds": ("WANVIDEOTEXTEMBEDS",),
"nag_scale": ("FLOAT", {"default": 11.0, "min": 0.0, "max": 100.0, "step": 0.1}),
"nag_tau": ("FLOAT", {"default": 2.5, "min": 0.0, "max": 10.0, "step": 0.1}),
"nag_alpha": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}),
},
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
RETURN_NAMES = ("text_embeds",)
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Adds NAG prompt embeds to original prompt embeds: 'https://github.com/ChenDarYen/Normalized-Attention-Guidance'"
def process(self, original_text_embeds, nag_text_embeds, nag_scale, nag_tau, nag_alpha):
prompt_embeds_dict_copy = original_text_embeds.copy()
prompt_embeds_dict_copy.update({
"nag_prompt_embeds": nag_text_embeds["prompt_embeds"],
"nag_params": {
"nag_scale": nag_scale,
"nag_tau": nag_tau,
"nag_alpha": nag_alpha,
}
})
return (prompt_embeds_dict_copy,)
class WanVideoTextEmbedBridge:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"positive": ("CONDITIONING",),
"negative": ("CONDITIONING",),
},
"optional": {
"negative": ("CONDITIONING",),
}
}
RETURN_TYPES = ("WANVIDEOTEXTEMBEDS", )
@@ -1292,11 +1371,11 @@ class WanVideoTextEmbedBridge:
CATEGORY = "WanVideoWrapper"
DESCRIPTION = "Bridge between ComfyUI native text embedding and WanVideoWrapper text embedding"
def process(self, positive, negative):
def process(self, positive, negative=None):
device=mm.get_torch_device()
prompt_embeds_dict = {
"prompt_embeds": positive[0][0].to(device),
"negative_prompt_embeds": negative[0][0].to(device),
"negative_prompt_embeds": negative[0][0].to(device) if negative is not None else None,
}
return (prompt_embeds_dict,)
@@ -2923,13 +3002,14 @@ class WanVideoSampler:
drift_timesteps = torch.cat([drift_timesteps, torch.tensor([0]).to(drift_timesteps.device)]).to(drift_timesteps.device)
timesteps[-drift_steps:] = drift_timesteps[-drift_steps:]
use_cfg_zero_star, use_fresca = False, False
use_cfg_zero_star = use_fresca = False
if experimental_args is not None:
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
if video_attention_split_steps:
transformer.video_attention_split_steps = [int(x.strip()) for x in video_attention_split_steps.split(",")]
else:
transformer.video_attention_split_steps = []
use_zero_init = experimental_args.get("use_zero_init", True)
use_cfg_zero_star = experimental_args.get("cfg_zero_star", False)
zero_star_steps = experimental_args.get("zero_star_steps", 0)
@@ -3051,12 +3131,18 @@ class WanVideoSampler:
"pcd_data": pcd_data,
"controlnet": controlnet,
"add_cond": add_cond_input,
"nag_params": text_embeds.get("nag_params", {}),
"nag_context": text_embeds.get("nag_prompt_embeds", None),
}
batch_size = 1
if not math.isclose(cfg_scale, 1.0) and len(positive_embeds) > 1:
negative_embeds = negative_embeds * len(positive_embeds)
# if use_nag:
# nag_negative_prompt_embeds = negative_embeds
# positive_embeds = torch.cat([positive_embeds[0], nag_negative_prompt_embeds], dim=0)
if not batched_cfg:
#cond
@@ -3440,7 +3526,7 @@ class WanVideoSampler:
noise_pred[:, c] += noise_pred_context * window_mask
counter[:, c] += window_mask
noise_pred /= counter
#normal inference
#region normal inference
else:
noise_pred, self.teacache_state = predict_with_cfg(
latent_model_input,
@@ -3725,6 +3811,7 @@ NODE_CLASS_MAPPINGS = {
"WanVideoSampler": WanVideoSampler,
"WanVideoDecode": WanVideoDecode,
"WanVideoTextEncode": WanVideoTextEncode,
"WanVideoTextEncodeSingle": WanVideoTextEncodeSingle,
"WanVideoModelLoader": WanVideoModelLoader,
"WanVideoVAELoader": WanVideoVAELoader,
"LoadWanVideoT5TextEncoder": LoadWanVideoT5TextEncoder,
@@ -3756,12 +3843,14 @@ NODE_CLASS_MAPPINGS = {
"WanVideoVACEModelSelect": WanVideoVACEModelSelect,
"WanVideoPhantomEmbeds": WanVideoPhantomEmbeds,
"CreateCFGScheduleFloatList": CreateCFGScheduleFloatList,
"WanVideoRealisDanceLatents": WanVideoRealisDanceLatents
"WanVideoRealisDanceLatents": WanVideoRealisDanceLatents,
"WanVideoApplyNAG": WanVideoApplyNAG
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoSampler": "WanVideo Sampler",
"WanVideoDecode": "WanVideo Decode",
"WanVideoTextEncode": "WanVideo TextEncode",
"WanVideoTextEncodeSingle": "WanVideo TextEncodeSingle",
"WanVideoTextImageEncode": "WanVideo TextImageEncode (IP2V)",
"WanVideoModelLoader": "WanVideo Model Loader",
"WanVideoVAELoader": "WanVideo VAE Loader",
@@ -3795,4 +3884,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoPhantomEmbeds": "WanVideo Phantom Embeds",
"CreateCFGScheduleFloatList": "WanVideo CFG Schedule Float List",
"WanVideoRealisDanceLatents": "WanVideo RealisDance Latents",
"WanVideoApplyNAG": "WanVideo Apply NAG"
}
+99 -57
View File
@@ -362,70 +362,95 @@ class WanSelfAttention(nn.Module):
x *= feta_scores
return x
def normalized_attention_guidance(self, b, n, d, q, context, nag_context=None, nag_params={}):
# NAG text attention
context_positive = context
context_negative = nag_context
nag_scale = nag_params['nag_scale']
nag_alpha = nag_params['nag_alpha']
nag_tau = nag_params['nag_tau']
k_positive = self.norm_k(self.k(context_positive)).view(b, -1, n, d)
v_positive = self.v(context_positive).view(b, -1, n, d)
k_negative = self.norm_k(self.k(context_negative)).view(b, -1, n, d)
v_negative = self.v(context_negative).view(b, -1, n, d)
x_positive = attention(q, k_positive, v_positive, k_lens=None, attention_mode=self.attention_mode)
x_positive = x_positive.flatten(2)
x_negative = attention(q, k_negative, v_negative, k_lens=None, attention_mode=self.attention_mode)
x_negative = x_negative.flatten(2)
nag_guidance = x_positive * nag_scale - x_negative * (nag_scale - 1)
norm_positive = torch.norm(x_positive, p=1, dim=-1, keepdim=True).expand_as(x_positive)
norm_guidance = torch.norm(nag_guidance, p=1, dim=-1, keepdim=True).expand_as(nag_guidance)
scale = norm_guidance / norm_positive
scale = torch.nan_to_num(scale, nan=10.0)
mask = scale > nag_tau
adjustment = (norm_positive * nag_tau) / (norm_guidance + 1e-7)
nag_guidance = torch.where(mask, nag_guidance * adjustment, nag_guidance)
return nag_guidance * nag_alpha + x_positive * (1 - nag_alpha)
#region T2V crossattn
class WanT2VCrossAttention(WanSelfAttention):
def forward(self, x, context, context_lens, clip_embed=None, audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
"""
def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(dim, num_heads, window_size, qk_norm, eps)
self.attention_mode = attention_mode
def forward(self, x, context, context_lens, clip_embed=None, audio_proj=None, audio_context_lens=None, audio_scale=1.0,
num_latent_frames=21, nag_params={}, nag_context=None):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
# compute query
q = self.norm_q(self.q(x)).view(b, -1, n, d)
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
# compute attention
x = attention(q, k, v, k_lens=context_lens, attention_mode=self.attention_mode)
if nag_context is not None:
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
else:
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, k_lens=None, attention_mode=self.attention_mode)
x_text = x_text.flatten(2)
# output
x = x.flatten(2)
x = x_text
# FantasyTalking audio attention
if audio_proj is not None:
if len(audio_proj.shape) == 4:
audio_q = q.view(b * num_latent_frames, -1, n, d) # [b, 21, l1, n, d]
audio_q = q.view(b * num_latent_frames, -1, n, d)
ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
audio_x = attention(
audio_q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode
)
audio_x = audio_x.view(b, q.size(1), n, d)
audio_x = audio_x.flatten(2)
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
elif len(audio_proj.shape) == 3:
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b, -1, n, d)
audio_x = attention(q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode)
audio_x = audio_x.flatten(2)
audio_x = attention(q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode).flatten(2)
x = x + audio_x * audio_scale
x = self.o(x)
return x
class WanI2VCrossAttention(WanSelfAttention):
def __init__(self,
dim,
num_heads,
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
attention_mode='sdpa'):
def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6, attention_mode='sdpa'):
super().__init__(dim, num_heads, window_size, qk_norm, eps)
self.k_img = nn.Linear(dim, dim)
self.v_img = nn.Linear(dim, dim)
# self.alpha = nn.Parameter(torch.zeros((1, )))
self.norm_k_img = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.attention_mode = attention_mode
def forward(self, x, context, context_lens, clip_embed, audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21):
def forward(self, x, context, context_lens, clip_embed, audio_proj=None, audio_context_lens=None,
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
@@ -433,41 +458,40 @@ class WanI2VCrossAttention(WanSelfAttention):
context_lens(Tensor): Shape [B]
"""
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
# compute query
q = self.norm_q(self.q(x)).view(b, -1, n, d)
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
# text attention
x = attention(q, k, v, k_lens=context_lens, attention_mode=self.attention_mode)
x = x.flatten(2)
if nag_context is not None:
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
else:
# text attention
k = self.norm_k(self.k(context)).view(b, -1, n, d)
v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, k_lens=context_lens, attention_mode=self.attention_mode).flatten(2)
#img attention
if clip_embed is not None:
k_img = self.norm_k_img(self.k_img(clip_embed)).view(b, -1, n, d)
v_img = self.v_img(clip_embed).view(b, -1, n, d)
img_x = attention(q, k_img, v_img, k_lens=None, attention_mode=self.attention_mode)
img_x = img_x.flatten(2)
x = x + img_x
img_x = attention(q, k_img, v_img, k_lens=None, attention_mode=self.attention_mode).flatten(2)
x = x_text + img_x
else:
x = x_text
# FantasyTalking audio attention
if audio_proj is not None:
if len(audio_proj.shape) == 4:
audio_q = q.view(b * num_latent_frames, -1, n, d) # [b, 21, l1, n, d]
audio_q = q.view(b * num_latent_frames, -1, n, d)
ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d)
audio_x = attention(
audio_q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode
)
audio_x = audio_x.view(b, q.size(1), n, d)
audio_x = audio_x.flatten(2)
audio_x = audio_x.view(b, q.size(1), n, d).flatten(2)
elif len(audio_proj.shape) == 3:
ip_key = self.k_proj(audio_proj).view(b, -1, n, d)
ip_value = self.v_proj(audio_proj).view(b, -1, n, d)
audio_x = attention(q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode)
audio_x = audio_x.flatten(2)
audio_x = attention(q, ip_key, ip_value, k_lens=audio_context_lens, attention_mode=self.attention_mode).flatten(2)
x = x + audio_x * audio_scale
@@ -538,6 +562,7 @@ class WanAttentionBlock(nn.Module):
def modulate(self, x, e):
return x * (1 + e[1]) + e[0]
#region attention forward
def forward(
self,
x,
@@ -556,8 +581,9 @@ class WanAttentionBlock(nn.Module):
audio_context_lens=None,
audio_scale=1.0,
num_latent_frames=21,
block_mask=None
block_mask=None,
nag_params={},
nag_context=None
):
r"""
Args:
@@ -595,7 +621,7 @@ class WanAttentionBlock(nn.Module):
input_x,
seq_lens, grid_sizes,
freqs, rope_func=rope_func,
block_mask=block_mask
block_mask=block_mask,
)
#ReCamMaster
if camera_embed is not None:
@@ -608,17 +634,21 @@ class WanAttentionBlock(nn.Module):
# cross-attention & ffn function
if (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1:
if nag_context is not None:
raise NotImplementedError("nag_context is not supported in split_cross_attn_ffn")
x = self.split_cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes)
else:
x = self.cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes,
audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale, num_latent_frames=num_latent_frames)
audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale,
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context)
del e
return x
@torch.compiler.disable()
def cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None,
audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21):
audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None):
x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed,
audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale, num_latent_frames=num_latent_frames)
audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale,
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context)
y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3])
x = x + (y * e[5])
return x
@@ -667,7 +697,7 @@ class WanAttentionBlock(nn.Module):
x_segment = x[:, segment_indices, :]
# Process segment with its prompt and clip embedding
processed_segment = self.cross_attn(self.norm3(x_segment), segment_context, segment_context_lens, clip_embed=segment_clip_embed)
processed_segment = self.cross_attn(self.norm3(x_segment), segment_context, segment_context_lens, clip_embed=segment_clip_embed, nag_scale=nag_scale)
processed_segment = processed_segment.to(x.dtype)
# Add to combined result
@@ -1164,8 +1194,9 @@ class WanModel(ModelMixin, ConfigMixin):
pcd_data=None,
controlnet=None,
add_cond=None,
attn_cond=None
attn_cond=None,
nag_params={},
nag_context=None
):
r"""
Forward pass through the diffusion model
@@ -1350,6 +1381,15 @@ class WanModel(ModelMixin, ConfigMixin):
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
for u in context
]).to(x.dtype))
# NAG
if nag_context is not None:
nag_context = self.text_embedding(
torch.stack([
torch.cat(
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
for u in nag_context
]).to(x.dtype))
if self.offload_txt_emb:
self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking)
@@ -1434,7 +1474,9 @@ class WanModel(ModelMixin, ConfigMixin):
audio_context_lens=audio_context_lens,
num_latent_frames = F,
audio_scale=audio_scale,
block_mask=self.block_mask
block_mask=self.block_mask,
nag_params=nag_params,
nag_context=nag_context
)
if vace_data is not None: