From 0ac366ba03eb06fa16626a770364c798ae8ea5bd Mon Sep 17 00:00:00 2001 From: kabachuha Date: Thu, 12 Jun 2025 14:43:51 +0300 Subject: [PATCH] add nag to attention internal parts --- nodes.py | 4 +- wanvideo/modules/model.py | 262 ++++++++++++++++++++++++++++++-------- 2 files changed, 210 insertions(+), 56 deletions(-) diff --git a/nodes.py b/nodes.py index d424923..95bcbf3 100644 --- a/nodes.py +++ b/nodes.py @@ -2320,6 +2320,7 @@ class WanVideoExperimentalArgs: "fresca_scale_low": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), "fresca_scale_high": ("FLOAT", {"default": 1.25, "min": 0.0, "max": 10.0, "step": 0.01}), "fresca_freq_cutoff": ("INT", {"default": 20, "min": 0, "max": 10000, "step": 1}), + "use_nag": ("BOOLEAN", {"default": False}), }, } @@ -2923,8 +2924,9 @@ 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, use_nag = False, False, False if experimental_args is not None: + use_nag = experimental_args.get("use_nag", False) 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(",")] diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index a574530..d6c6499 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -363,117 +363,269 @@ class WanSelfAttention(nn.Module): return x - class WanT2VCrossAttention(WanSelfAttention): + def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6, attention_mode='sdpa', + nag_scale=1.0, nag_tau=2.5, nag_alpha=0.25): + super().__init__(dim, num_heads, window_size, qk_norm, eps) + self.attention_mode = attention_mode + self.nag_scale = nag_scale + 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): - r""" - Args: - x(Tensor): Shape [B, L1, C] - context(Tensor): Shape [B, L2, C] - context_lens(Tensor): Shape [B] - """ b, n, d = x.size(0), self.num_heads, self.head_dim - - # compute query, key, value 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) + apply_guidance = self.nag_scale > 1 and context is not None + if apply_guidance and context.size(0) == 2 * b: + batch_size = b + context_positive = context[:batch_size] + context_negative = context[batch_size:] + context_lens_positive = context_lens[:batch_size] + context_lens_negative = context_lens[batch_size:] - # output - x = x.flatten(2) + k_positive = self.norm_k(self.k(context_positive)).view(batch_size, -1, n, d) + v_positive = self.v(context_positive).view(batch_size, -1, n, d) + k_negative = self.norm_k(self.k(context_negative)).view(batch_size, -1, n, d) + v_negative = self.v(context_negative).view(batch_size, -1, n, d) + + hidden_states_positive = attention(q, k_positive, v_positive, k_lens=context_lens_positive, attention_mode=self.attention_mode) + hidden_states_positive = hidden_states_positive.flatten(2) + + hidden_states_negative = attention(q, k_negative, v_negative, k_lens=context_lens_negative, attention_mode=self.attention_mode) + hidden_states_negative = hidden_states_negative.flatten(2) + + hidden_states_guidance = hidden_states_positive * self.nag_scale - hidden_states_negative * (self.nag_scale - 1) + + norm_positive = torch.norm(hidden_states_positive, p=1, dim=-1, keepdim=True).expand_as(hidden_states_positive) + norm_guidance = torch.norm(hidden_states_guidance, p=1, dim=-1, keepdim=True).expand_as(hidden_states_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) + hidden_states_guidance = torch.where(mask, hidden_states_guidance * adjustment, hidden_states_guidance) + + x_text = hidden_states_guidance * self.nag_alpha + hidden_states_positive * (1 - self.nag_alpha) + 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=context_lens, attention_mode=self.attention_mode) + x_text = x_text.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 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] +# """ +# b, n, d = x.size(0), self.num_heads, self.head_dim + +# # compute query, key, value +# 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) + +# # output +# x = x.flatten(2) + +# # 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] +# 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) +# 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) + +# 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', + nag_scale=1.0, nag_tau=2.5, nag_alpha=0.25): 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 + self.nag_scale = nag_scale + 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): - r""" - Args: - x(Tensor): Shape [B, L1, C] - context(Tensor): Shape [B, L2, C] - context_lens(Tensor): Shape [B] - """ b, n, d = x.size(0), self.num_heads, self.head_dim - - # compute query, key, value 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) - #img attention + apply_guidance = self.nag_scale > 1 and context is not None + if apply_guidance and context.size(0) == 2 * b: + batch_size = b + context_positive = context[:batch_size] + context_negative = context[batch_size:] + context_lens_positive = context_lens[:batch_size] + context_lens_negative = context_lens[batch_size:] + + k_positive = self.norm_k(self.k(context_positive)).view(batch_size, -1, n, d) + v_positive = self.v(context_positive).view(batch_size, -1, n, d) + k_negative = self.norm_k(self.k(context_negative)).view(batch_size, -1, n, d) + v_negative = self.v(context_negative).view(batch_size, -1, n, d) + + hidden_states_positive = attention(q, k_positive, v_positive, k_lens=context_lens_positive, attention_mode=self.attention_mode) + hidden_states_positive = hidden_states_positive.flatten(2) + + hidden_states_negative = attention(q, k_negative, v_negative, k_lens=context_lens_negative, attention_mode=self.attention_mode) + hidden_states_negative = hidden_states_negative.flatten(2) + + hidden_states_guidance = hidden_states_positive * self.nag_scale - hidden_states_negative * (self.nag_scale - 1) + + norm_positive = torch.norm(hidden_states_positive, p=1, dim=-1, keepdim=True).expand_as(hidden_states_positive) + norm_guidance = torch.norm(hidden_states_guidance, p=1, dim=-1, keepdim=True).expand_as(hidden_states_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) + hidden_states_guidance = torch.where(mask, hidden_states_guidance * adjustment, hidden_states_guidance) + + x_text = hidden_states_guidance * self.nag_alpha + hidden_states_positive * (1 - self.nag_alpha) + 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=context_lens, attention_mode=self.attention_mode).flatten(2) + 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) + 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 - x = x + img_x - - # 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'): +# 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): +# r""" +# Args: +# x(Tensor): Shape [B, L1, C] +# context(Tensor): Shape [B, L2, C] +# context_lens(Tensor): Shape [B] +# """ +# b, n, d = x.size(0), self.num_heads, self.head_dim + +# # compute query, key, value +# 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) + +# #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 + +# # 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] +# 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) +# 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) + +# x = x + audio_x * audio_scale + +# x = self.o(x) +# return x + WAN_CROSSATTENTION_CLASSES = { 't2v_cross_attn': WanT2VCrossAttention,