add nag to attention internal parts
This commit is contained in:
@@ -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(",")]
|
||||
|
||||
+205
-53
@@ -363,117 +363,269 @@ class WanSelfAttention(nn.Module):
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
# compute query, key, value
|
||||
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):
|
||||
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
|
||||
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)
|
||||
x_text = x_text.flatten(2)
|
||||
|
||||
# compute attention
|
||||
x = attention(q, k, v, k_lens=context_lens, attention_mode=self.attention_mode)
|
||||
x = x_text
|
||||
|
||||
# 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]
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
# 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)
|
||||
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,
|
||||
|
||||
Reference in New Issue
Block a user