Cleanup, create own node for NAG
This commit is contained in:
@@ -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,)
|
||||
|
||||
@@ -3060,8 +3139,8 @@ class WanVideoSampler:
|
||||
"pcd_data": pcd_data,
|
||||
"controlnet": controlnet,
|
||||
"add_cond": add_cond_input,
|
||||
"nag_scale": nag_scale,
|
||||
"nag_context": nag_negative_context
|
||||
"nag_params": text_embeds.get("nag_params", {}),
|
||||
"nag_context": text_embeds.get("nag_prompt_embeds", None),
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -3455,7 +3534,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,
|
||||
@@ -3740,6 +3819,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoSampler": WanVideoSampler,
|
||||
"WanVideoDecode": WanVideoDecode,
|
||||
"WanVideoTextEncode": WanVideoTextEncode,
|
||||
"WanVideoTextEncodeSingle": WanVideoTextEncodeSingle,
|
||||
"WanVideoModelLoader": WanVideoModelLoader,
|
||||
"WanVideoVAELoader": WanVideoVAELoader,
|
||||
"LoadWanVideoT5TextEncoder": LoadWanVideoT5TextEncoder,
|
||||
@@ -3771,12 +3851,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",
|
||||
@@ -3810,4 +3892,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoPhantomEmbeds": "WanVideo Phantom Embeds",
|
||||
"CreateCFGScheduleFloatList": "WanVideo CFG Schedule Float List",
|
||||
"WanVideoRealisDanceLatents": "WanVideo RealisDance Latents",
|
||||
"WanVideoApplyNAG": "WanVideo Apply NAG"
|
||||
}
|
||||
|
||||
+53
-75
@@ -362,50 +362,55 @@ 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 __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6, attention_mode='sdpa', nag_tau=2.5, nag_alpha=0.25):
|
||||
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
|
||||
self.nag_tau = nag_tau
|
||||
self.nag_alpha = nag_alpha
|
||||
|
||||
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_scale=11.0, nag_context=None):
|
||||
num_latent_frames=21, nag_params={}, nag_context=None):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
# compute query
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
|
||||
if nag_scale > 1 and nag_context is not None:
|
||||
# NAG text attention
|
||||
context_positive = context
|
||||
context_negative = nag_context
|
||||
|
||||
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 > self.nag_tau
|
||||
adjustment = (norm_positive * self.nag_tau) / (norm_guidance + 1e-7)
|
||||
nag_guidance = torch.where(mask, nag_guidance * adjustment, nag_guidance)
|
||||
|
||||
x_text = nag_guidance * self.nag_alpha + x_positive * (1 - self.nag_alpha)
|
||||
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)
|
||||
@@ -437,17 +442,15 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
|
||||
class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6, attention_mode='sdpa', nag_tau=2.5, nag_alpha=0.25):
|
||||
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.norm_k_img = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.attention_mode = attention_mode
|
||||
self.nag_tau = nag_tau
|
||||
self.nag_alpha = nag_alpha
|
||||
|
||||
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_scale=11.0, nag_context=None):
|
||||
audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L1, C]
|
||||
@@ -458,35 +461,8 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
# compute query
|
||||
q = self.norm_q(self.q(x)).view(b, -1, n, d)
|
||||
|
||||
if nag_scale > 1 and nag_context is not None:
|
||||
# NAG text attention
|
||||
context_positive = context
|
||||
context_negative = nag_context
|
||||
|
||||
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 > self.nag_tau
|
||||
adjustment = (norm_positive * self.nag_tau) / (norm_guidance + 1e-7)
|
||||
nag_guidance = torch.where(mask, nag_guidance * adjustment, nag_guidance)
|
||||
|
||||
x_text = nag_guidance * self.nag_alpha + x_positive * (1 - self.nag_alpha)
|
||||
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)
|
||||
@@ -606,7 +582,7 @@ class WanAttentionBlock(nn.Module):
|
||||
audio_scale=1.0,
|
||||
num_latent_frames=21,
|
||||
block_mask=None,
|
||||
nag_scale=11.0,
|
||||
nag_params={},
|
||||
nag_context=None
|
||||
):
|
||||
r"""
|
||||
@@ -657,26 +633,28 @@ class WanAttentionBlock(nn.Module):
|
||||
del y
|
||||
|
||||
# cross-attention & ffn function
|
||||
if (((context.shape[0] // 2 > 1) if nag_scale > 1 else (context.shape[0] > 1)) or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1:
|
||||
x = self.split_cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes, nag_scale=nag_scale)
|
||||
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, nag_scale=nag_scale, nag_context=nag_context)
|
||||
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, nag_scale=11.0, nag_context=None):
|
||||
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, nag_scale=nag_scale, nag_context=nag_context)
|
||||
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
|
||||
|
||||
@torch.compiler.disable()
|
||||
def split_cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None, nag_scale=11.):
|
||||
def split_cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None):
|
||||
# Get number of prompts
|
||||
num_prompts = context.shape[0]
|
||||
num_clip_embeds = 0 if clip_embed is None else clip_embed.shape[0]
|
||||
@@ -1217,7 +1195,7 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
controlnet=None,
|
||||
add_cond=None,
|
||||
attn_cond=None,
|
||||
nag_scale=11.,
|
||||
nag_params={},
|
||||
nag_context=None
|
||||
):
|
||||
r"""
|
||||
@@ -1497,7 +1475,7 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
num_latent_frames = F,
|
||||
audio_scale=audio_scale,
|
||||
block_mask=self.block_mask,
|
||||
nag_scale=nag_scale,
|
||||
nag_params=nag_params,
|
||||
nag_context=nag_context
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user