From bdf7a3d02ac62754dc4181afb6a1228f2946db26 Mon Sep 17 00:00:00 2001 From: kabachuha Date: Thu, 12 Jun 2025 15:26:59 +0300 Subject: [PATCH] pass nag scale into the model --- nodes.py | 6 +++++ wanvideo/modules/model.py | 49 +++++++++++++++++++-------------------- 2 files changed, 30 insertions(+), 25 deletions(-) diff --git a/nodes.py b/nodes.py index d359722..4b5647c 100644 --- a/nodes.py +++ b/nodes.py @@ -2321,6 +2321,7 @@ class WanVideoExperimentalArgs: "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}), + "nag_scale": ("FLOAT", {"default": 11.0, "min": 1.0, "max": 20.0, "step": 1.0}), }, } @@ -2927,6 +2928,7 @@ class WanVideoSampler: 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) + nag_scale = experimental_args.get("nag_scale", 11) 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(",")] @@ -3053,6 +3055,7 @@ class WanVideoSampler: "pcd_data": pcd_data, "controlnet": controlnet, "add_cond": add_cond_input, + "nag_scale": nag_scale, } batch_size = 1 @@ -3061,8 +3064,11 @@ class WanVideoSampler: negative_embeds = negative_embeds * len(positive_embeds) if use_nag: + print("Nag triggered!") nag_negative_prompt_embeds = negative_embeds + print(positive_embeds.shape) positive_embeds = torch.cat([positive_embeds, nag_negative_prompt_embeds], dim=0) + print(positive_embeds.shape) if not batched_cfg: #cond diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index d6c6499..19f8526 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -365,19 +365,17 @@ class WanSelfAttention(nn.Module): 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): + 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): 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): + 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): b, n, d = x.size(0), self.num_heads, self.head_dim q = self.norm_q(self.q(x)).view(b, -1, n, d) - apply_guidance = self.nag_scale > 1 and context is not None + apply_guidance = 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] @@ -396,7 +394,7 @@ class WanT2VCrossAttention(WanSelfAttention): 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) + hidden_states_guidance = hidden_states_positive * nag_scale - hidden_states_negative * (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) @@ -409,6 +407,8 @@ class WanT2VCrossAttention(WanSelfAttention): 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) + + print("Nag applied in apply_guidance!") else: k = self.norm_k(self.k(context)).view(b, -1, n, d) v = self.v(context).view(b, -1, n, d) @@ -481,22 +481,20 @@ 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_scale=1.0, 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', 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.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): + 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.): b, n, d = x.size(0), self.num_heads, self.head_dim q = self.norm_q(self.q(x)).view(b, -1, n, d) - apply_guidance = self.nag_scale > 1 and context is not None + apply_guidance = 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] @@ -515,7 +513,7 @@ class WanI2VCrossAttention(WanSelfAttention): 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) + hidden_states_guidance = hidden_states_positive * nag_scale - hidden_states_negative * (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) @@ -708,8 +706,8 @@ class WanAttentionBlock(nn.Module): audio_context_lens=None, audio_scale=1.0, num_latent_frames=21, - block_mask=None - + block_mask=None, + nag_scale=11. ): r""" Args: @@ -747,7 +745,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: @@ -759,24 +757,24 @@ class WanAttentionBlock(nn.Module): del y # 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: - x = self.split_cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) + 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) 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_scale=nag_scale) 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_scale=11.): 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_scale=nag_scale) 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): + def split_cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None, nag_scale=11.): # Get number of prompts num_prompts = context.shape[0] num_clip_embeds = 0 if clip_embed is None else clip_embed.shape[0] @@ -819,7 +817,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 @@ -1316,8 +1314,8 @@ class WanModel(ModelMixin, ConfigMixin): pcd_data=None, controlnet=None, add_cond=None, - attn_cond=None - + attn_cond=None, + nag_scale=11., ): r""" Forward pass through the diffusion model @@ -1586,7 +1584,8 @@ 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_scale=nag_scale ) if vace_data is not None: