Merge branch 'main' into dev
This commit is contained in:
+24
-78
@@ -1,11 +1,8 @@
|
||||
from diffusers import ModelMixin, ConfigMixin
|
||||
from einops import rearrange, repeat
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from ..wanvideo.modules.attention import attention
|
||||
|
||||
from comfy import model_management as mm
|
||||
|
||||
def timestep_transform(
|
||||
t,
|
||||
shift=5.0,
|
||||
@@ -65,7 +62,6 @@ def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', att
|
||||
x_ref_attn_map_source = x_ref_attn_map_source.to(visual_q.dtype)
|
||||
|
||||
for class_idx, ref_target_mask in enumerate(ref_target_masks):
|
||||
mm.soft_empty_cache()
|
||||
ref_target_mask = ref_target_mask[None, None, None, ...]
|
||||
x_ref_attnmap = x_ref_attn_map_source * ref_target_mask
|
||||
x_ref_attnmap = x_ref_attnmap.sum(-1) / ref_target_mask.sum() # B, H, x_seqlens, ref_seqlens --> B, H, x_seqlens
|
||||
@@ -78,13 +74,11 @@ def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', att
|
||||
|
||||
x_ref_attn_maps.append(x_ref_attnmap)
|
||||
|
||||
del attn
|
||||
del x_ref_attn_map_source
|
||||
mm.soft_empty_cache()
|
||||
del attn, x_ref_attn_map_source
|
||||
|
||||
return torch.concat(x_ref_attn_maps, dim=0)
|
||||
|
||||
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2, enable_sp=False):
|
||||
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2):
|
||||
"""Args:
|
||||
query (torch.tensor): B M H K
|
||||
key (torch.tensor): B M H K
|
||||
@@ -145,7 +139,7 @@ class RotaryPositionalEmbedding1D(nn.Module):
|
||||
|
||||
return x_.type_as(x)
|
||||
|
||||
class AudioProjModel(ModelMixin, ConfigMixin):
|
||||
class AudioProjModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
seq_len=5,
|
||||
@@ -217,11 +211,6 @@ class SingleStreamAttention(nn.Module):
|
||||
encoder_hidden_states_dim: int,
|
||||
num_heads: int,
|
||||
qkv_bias: bool,
|
||||
qk_norm: bool,
|
||||
norm_layer: nn.Module,
|
||||
attn_drop: float = 0.0,
|
||||
proj_drop: float = 0.0,
|
||||
eps: float = 1e-6,
|
||||
attention_mode: str = 'sdpa',
|
||||
) -> None:
|
||||
super().__init__()
|
||||
@@ -230,74 +219,46 @@ class SingleStreamAttention(nn.Module):
|
||||
self.encoder_hidden_states_dim = encoder_hidden_states_dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.scale = self.head_dim**-0.5
|
||||
self.qk_norm = qk_norm
|
||||
|
||||
self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
|
||||
self.q_norm = norm_layer(self.head_dim, eps=eps) if qk_norm else nn.Identity()
|
||||
self.k_norm = norm_layer(self.head_dim,eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias)
|
||||
|
||||
self.add_q_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
|
||||
self.add_k_norm = norm_layer(self.head_dim) if qk_norm else nn.Identity()
|
||||
|
||||
self.attention_mode = attention_mode
|
||||
|
||||
def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None, enable_sp=False, kv_seq=None) -> torch.Tensor:
|
||||
self.q_linear = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.kv_linear = nn.Linear(encoder_hidden_states_dim, dim * 2, bias=qkv_bias)
|
||||
|
||||
def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, shape=None) -> torch.Tensor:
|
||||
N_t, N_h, N_w = shape
|
||||
|
||||
expected_tokens = N_t * N_h * N_w
|
||||
actual_tokens = x.shape[1]
|
||||
x_extra = None
|
||||
if x.shape[0] * N_t != encoder_hidden_states.shape[0]:
|
||||
|
||||
if actual_tokens != expected_tokens:
|
||||
x_extra = x[:, -N_h * N_w:, :]
|
||||
x = x[:, :-N_h * N_w, :]
|
||||
N_t = N_t - 1
|
||||
x = rearrange(x, "B (N_t S) C -> (B N_t) S C", N_t=N_t)
|
||||
|
||||
B = x.shape[0]
|
||||
S = N_h * N_w
|
||||
x = x.view(B * N_t, S, self.dim)
|
||||
|
||||
# get q for hidden_state
|
||||
B, N, C = x.shape
|
||||
q = self.q_linear(x)
|
||||
q_shape = (B, N, self.num_heads, self.head_dim)
|
||||
q = q.view(q_shape).permute((0, 2, 1, 3))
|
||||
|
||||
if self.qk_norm:
|
||||
q = self.q_norm(q)
|
||||
q = self.q_linear(x).view(B * N_t, S, self.num_heads, self.head_dim)
|
||||
|
||||
# get kv from encoder_hidden_states
|
||||
_, N_a, _ = encoder_hidden_states.shape
|
||||
encoder_kv = self.kv_linear(encoder_hidden_states)
|
||||
encoder_kv_shape = (B, N_a, 2, self.num_heads, self.head_dim)
|
||||
encoder_kv = encoder_kv.view(encoder_kv_shape).permute((2, 0, 3, 1, 4))
|
||||
encoder_k, encoder_v = encoder_kv.unbind(0)
|
||||
# get kv from encoder_hidden_states # shape: (B, N, num_heads, head_dim)
|
||||
kv = self.kv_linear(encoder_hidden_states)
|
||||
encoder_k, encoder_v = kv.view(B * N_t, encoder_hidden_states.shape[1], 2, self.num_heads, self.head_dim).unbind(2)
|
||||
|
||||
if self.qk_norm:
|
||||
encoder_k = self.add_k_norm(encoder_k)
|
||||
|
||||
x = attention(
|
||||
q.transpose(1, 2),
|
||||
encoder_k.transpose(1, 2),
|
||||
encoder_v.transpose(1, 2),
|
||||
attention_mode=self.attention_mode
|
||||
)
|
||||
x = attention(q, encoder_k, encoder_v, attention_mode=self.attention_mode)
|
||||
|
||||
# linear transform
|
||||
x_output_shape = (B, N, C)
|
||||
#x = x.transpose(1, 2)
|
||||
x = x.reshape(x_output_shape)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
|
||||
x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t)
|
||||
x = self.proj(x.reshape(B * N_t, S, self.dim))
|
||||
x = x.view(B, N_t * S, self.dim)
|
||||
|
||||
if x_extra is not None:
|
||||
x = torch.cat([x, torch.zeros_like(x_extra)], dim=1)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
"""Multi-speaker rotary-position cross-attention.
|
||||
@@ -315,11 +276,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
encoder_hidden_states_dim: int,
|
||||
num_heads: int,
|
||||
qkv_bias: bool,
|
||||
qk_norm: bool,
|
||||
norm_layer: nn.Module,
|
||||
attn_drop: float = 0.0,
|
||||
proj_drop: float = 0.0,
|
||||
eps: float = 1e-6,
|
||||
class_range: int = 24,
|
||||
class_interval: int = 4,
|
||||
attention_mode: str = 'sdpa',
|
||||
@@ -329,11 +285,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
encoder_hidden_states_dim=encoder_hidden_states_dim,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_norm=qk_norm,
|
||||
norm_layer=norm_layer,
|
||||
attn_drop=attn_drop,
|
||||
proj_drop=proj_drop,
|
||||
eps=eps,
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
|
||||
@@ -376,8 +327,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
B, N, C = x.shape
|
||||
q = self.q_linear(x)
|
||||
q = q.view(B, N, self.num_heads, self.head_dim).permute(0, 2, 1, 3)
|
||||
if self.qk_norm:
|
||||
q = self.q_norm(q)
|
||||
|
||||
if human_num == 2:
|
||||
# Use `class_range` logic for exactly 2 speakers
|
||||
@@ -441,8 +390,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
encoder_kv = self.kv_linear(encoder_hidden_states)
|
||||
encoder_kv = encoder_kv.view(B, N_a, 2, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4)
|
||||
encoder_k, encoder_v = encoder_kv.unbind(0)
|
||||
if self.qk_norm:
|
||||
encoder_k = self.add_k_norm(encoder_k)
|
||||
|
||||
# Rotary for keys – assign centre of each speaker bucket to its context tokens
|
||||
if human_num == 2:
|
||||
@@ -478,7 +425,6 @@ class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
# Linear projection
|
||||
x = x.reshape(B, N, C)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
|
||||
# Restore original layout
|
||||
x = rearrange(x, "(B N_t) S C -> B (N_t S) C", N_t=N_t)
|
||||
|
||||
Reference in New Issue
Block a user