pass nag scale into the model
This commit is contained in:
@@ -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
|
||||
|
||||
+24
-25
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user