add nag to attention internal parts

This commit is contained in:
kabachuha
2025-06-12 14:43:51 +03:00
parent 6eddec54a6
commit 0ac366ba03
2 changed files with 210 additions and 56 deletions
+3 -1
View File
@@ -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(",")]
+207 -55
View File
@@ -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,