VRAM optimizations
Minor for almost everything, major for multitalk when using masks
This commit is contained in:
+40
-35
@@ -42,41 +42,43 @@ def rotate_half(x):
|
||||
x = torch.stack((-x2, x1), dim=-1)
|
||||
return rearrange(x, "... d r -> ... (d r)")
|
||||
|
||||
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, mode='mean', attn_bias=None):
|
||||
|
||||
ref_k = ref_k.to(visual_q.dtype).to(visual_q.device)
|
||||
def calculate_x_ref_attn_map(visual_q, ref_k, ref_target_masks, split_num=4):
|
||||
scale = 1.0 / visual_q.shape[-1] ** 0.5
|
||||
visual_q = visual_q * scale
|
||||
visual_q = visual_q.transpose(1, 2)
|
||||
ref_k = ref_k.transpose(1, 2)
|
||||
attn = visual_q @ ref_k.transpose(-2, -1)
|
||||
|
||||
if attn_bias is not None:
|
||||
attn = attn + attn_bias
|
||||
|
||||
x_ref_attn_map_source = attn.softmax(-1) # B, H, x_seqlens, ref_seqlens
|
||||
visual_q = visual_q.transpose(1, 2) * scale
|
||||
|
||||
B, H, x_seqlens, K = visual_q.shape
|
||||
|
||||
x_ref_attn_maps = []
|
||||
ref_target_masks = ref_target_masks.to(visual_q.dtype)
|
||||
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):
|
||||
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
|
||||
x_ref_attnmap = x_ref_attnmap.permute(0, 2, 1) # B, x_seqlens, H
|
||||
|
||||
if mode == 'mean':
|
||||
x_ref_attnmap = x_ref_attnmap.mean(-1) # B, x_seqlens
|
||||
elif mode == 'max':
|
||||
x_ref_attnmap = x_ref_attnmap.max(-1) # B, x_seqlens
|
||||
|
||||
x_ref_attn_maps.append(x_ref_attnmap)
|
||||
|
||||
del attn, x_ref_attn_map_source
|
||||
ref_target_mask = ref_target_mask.view(1, 1, 1, -1)
|
||||
|
||||
return torch.concat(x_ref_attn_maps, dim=0)
|
||||
x_ref_attnmap = torch.zeros(B, H, x_seqlens, device=visual_q.device, dtype=visual_q.dtype)
|
||||
chunk_size = min(max(x_seqlens // split_num, 1), x_seqlens)
|
||||
|
||||
for i in range(0, x_seqlens, chunk_size):
|
||||
end_i = min(i + chunk_size, x_seqlens)
|
||||
|
||||
attn_chunk = visual_q[:, :, i:end_i] @ ref_k.permute(0, 2, 3, 1) # B, H, chunk, ref_seqlens
|
||||
|
||||
# Apply softmax
|
||||
attn_max = attn_chunk.max(dim=-1, keepdim=True).values
|
||||
attn_chunk = (attn_chunk - attn_max).exp()
|
||||
attn_sum = attn_chunk.sum(dim=-1, keepdim=True)
|
||||
attn_chunk = attn_chunk / (attn_sum + 1e-8)
|
||||
|
||||
# Apply mask and sum
|
||||
masked_attn = attn_chunk * ref_target_mask
|
||||
x_ref_attnmap[:, :, i:end_i] = masked_attn.sum(-1) / (ref_target_mask.sum() + 1e-8)
|
||||
|
||||
del attn_chunk, masked_attn
|
||||
|
||||
# Average across heads
|
||||
x_ref_attnmap = x_ref_attnmap.mean(dim=1) # B, x_seqlens
|
||||
x_ref_attn_maps.append(x_ref_attnmap)
|
||||
|
||||
del visual_q, ref_k
|
||||
|
||||
return torch.cat(x_ref_attn_maps, dim=0)
|
||||
|
||||
def get_attn_map_with_target(visual_q, ref_k, shape, ref_target_masks=None, split_num=2):
|
||||
"""Args:
|
||||
@@ -129,15 +131,18 @@ class RotaryPositionalEmbedding1D(nn.Module):
|
||||
query with the same shape as input.
|
||||
"""
|
||||
freqs_cis = self.precompute_freqs_cis_1d(pos_indices)
|
||||
|
||||
x_ = x.float()
|
||||
in_dtype = x.dtype
|
||||
x = x.float()
|
||||
|
||||
freqs_cis = freqs_cis.float().to(x.device)
|
||||
cos, sin = freqs_cis.cos(), freqs_cis.sin()
|
||||
cos, sin = rearrange(cos, 'n d -> 1 1 n d'), rearrange(sin, 'n d -> 1 1 n d')
|
||||
x_ = (x_ * cos) + (rotate_half(x_) * sin)
|
||||
cos = rearrange(freqs_cis.cos(), 'n d -> 1 1 n d')
|
||||
sin = rearrange(freqs_cis.sin(), 'n d -> 1 1 n d')
|
||||
|
||||
return x_.type_as(x)
|
||||
# In-place rotation to save memory
|
||||
x_rotated = rotate_half(x)
|
||||
x.mul_(cos).add_(x_rotated * sin)
|
||||
|
||||
return x.to(in_dtype)
|
||||
|
||||
class AudioProjModel(nn.Module):
|
||||
def __init__(
|
||||
|
||||
@@ -764,7 +764,8 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype)
|
||||
v_img = self.v_img(clip_embed).view(b, -1, n, d)
|
||||
img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
|
||||
x = x_text + img_x
|
||||
x_text.add_(img_x)
|
||||
x = x_text
|
||||
else:
|
||||
x = x_text
|
||||
|
||||
@@ -1292,7 +1293,7 @@ class WanAttentionBlock(nn.Module):
|
||||
y[:, tr_end:] * gate_msa
|
||||
], dim=1).to(input_dtype)
|
||||
else:
|
||||
x = x.addcmul(y, gate_msa)
|
||||
x.addcmul_(y, gate_msa)
|
||||
del y, gate_msa
|
||||
|
||||
# cross-attention & ffn function
|
||||
@@ -1322,11 +1323,10 @@ class WanAttentionBlock(nn.Module):
|
||||
x = self.split_cross_attn_ffn(x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed, grid_sizes)
|
||||
return x, x_ip, lynx_ref_feature, x_ovi
|
||||
else:
|
||||
x = x + self.cross_attn(self.norm3(x.to(self.norm3.weight.dtype)).to(input_dtype), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
x += self.cross_attn(self.norm3(x.to(self.norm3.weight.dtype)).to(input_dtype), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale,
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context,
|
||||
rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs,
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, longcat_num_cond_latents=longcat_num_cond_latents)
|
||||
x = x.to(input_dtype)
|
||||
adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, longcat_num_cond_latents=longcat_num_cond_latents).to(input_dtype)
|
||||
# MultiTalk
|
||||
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
|
||||
|
||||
@@ -1340,7 +1340,8 @@ class WanAttentionBlock(nn.Module):
|
||||
else:
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding,
|
||||
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
|
||||
x = x.add(x_audio, alpha=audio_scale)
|
||||
x.add_(x_audio, alpha=audio_scale)
|
||||
del x_audio
|
||||
|
||||
# MTV-Crafter Motion Attention
|
||||
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
|
||||
@@ -2577,6 +2578,7 @@ class WanModel(torch.nn.Module):
|
||||
scail_x = [u.flatten(2).transpose(1, 2) * scail_input.get("pose_strength", 1) for u in scail_x]
|
||||
x = [torch.cat([u, v], dim=1) for u, v in zip(x, scail_x)]
|
||||
seq_len += scail_x[0].shape[1]
|
||||
del scail_x
|
||||
pose_frame_shape = scail_pose_latents.shape
|
||||
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
|
||||
|
||||
Reference in New Issue
Block a user