small possible optimizations
This commit is contained in:
+61
-51
@@ -201,15 +201,10 @@ class WanSelfAttention(nn.Module):
|
||||
"""
|
||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||
|
||||
# query, key, value function
|
||||
def qkv_fn(x):
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
return q, k, v
|
||||
|
||||
q, k, v = qkv_fn(x)
|
||||
|
||||
# query, key, value
|
||||
q = self.norm_q(self.q(x)).view(b, s, n, d)
|
||||
k = self.norm_k(self.k(x)).view(b, s, n, d)
|
||||
v = self.v(x).view(b, s, n, d)
|
||||
|
||||
if is_enhance_enabled():
|
||||
feta_scores = get_feta_scores(q, k)
|
||||
@@ -573,7 +568,7 @@ class WanAttentionBlock(nn.Module):
|
||||
return e
|
||||
|
||||
def modulate(self, x, e):
|
||||
return x * (1 + e[1]) + e[0]
|
||||
return x.mul_(1 + e[1]).add_(e[0])
|
||||
|
||||
#region attention forward
|
||||
def forward(
|
||||
@@ -647,7 +642,7 @@ class WanAttentionBlock(nn.Module):
|
||||
|
||||
del input_x
|
||||
|
||||
x = x + (y * e[2])
|
||||
x.add_(y.mul_(e[2]))
|
||||
del y
|
||||
|
||||
# cross-attention & ffn function
|
||||
@@ -662,8 +657,10 @@ class WanAttentionBlock(nn.Module):
|
||||
num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond,
|
||||
multitalk_audio_embedding=multitalk_audio_embedding, x_ref_attn_map=x_ref_attn_map, human_num=human_num)
|
||||
else:
|
||||
y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3])
|
||||
x = x + (y * e[5])
|
||||
y = self.ffn(self.norm2(x).mul_(1 + e[4]).add_(e[3]))
|
||||
x.add_(y.mul_(e[5]))
|
||||
del y
|
||||
|
||||
|
||||
del e
|
||||
return x
|
||||
@@ -679,8 +676,9 @@ class WanAttentionBlock(nn.Module):
|
||||
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
|
||||
shape=grid_sizes[0], x_ref_attn_map=x_ref_attn_map, human_num=human_num)
|
||||
x = x + x_audio * audio_scale
|
||||
y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3])
|
||||
x = x + (y * e[5])
|
||||
y = self.ffn(self.norm2(x).mul_(1 + e[4]).add_(e[3]))
|
||||
x.add_(y.mul_(e[5]))
|
||||
del y
|
||||
return x
|
||||
|
||||
@torch.compiler.disable()
|
||||
@@ -842,7 +840,7 @@ class Head(nn.Module):
|
||||
# x = self.head(normed * (1 + e[1]) + e[0])
|
||||
|
||||
e = self.get_mod(e)
|
||||
x = self.head(self.norm(x) * (1 + e[1]) + e[0])
|
||||
x = self.head(self.norm(x).mul_(1 + e[1]).add_(e[0]))
|
||||
return x
|
||||
|
||||
|
||||
@@ -1071,6 +1069,7 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
num_frames=None,
|
||||
k=None,
|
||||
)
|
||||
self.cached_freqs = self.cached_shape = self.cached_cond = None
|
||||
|
||||
# buffers (don't use register_buffer otherwise dtype will be changed in to())
|
||||
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
||||
@@ -1356,41 +1355,52 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
])
|
||||
|
||||
if freqs is None: #comfy rope
|
||||
rope_func = "comfy"
|
||||
f_len = ((F + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||
h_len = ((H + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||
w_len = ((W + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||
img_ids = torch.zeros((f_len, h_len, w_len, 3), device=x.device, dtype=x.dtype)
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(0, f_len - 1, steps=f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(0, h_len - 1, steps=h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(0, w_len - 1, steps=w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
|
||||
if attn_cond is not None:
|
||||
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=x.device, dtype=x.dtype)
|
||||
|
||||
#shift
|
||||
shift_f_size = 81 # Default value
|
||||
shift_f = False
|
||||
if shift_f:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(shift_f_size, shift_f_size + cond_f_len - 1,steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
else:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
|
||||
# Combine original and conditional position ids
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
cond_img_ids = repeat(cond_img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
|
||||
|
||||
# Generate RoPE frequencies for the combined positions
|
||||
freqs = self.rope_embedder(combined_img_ids).movedim(1, 2)
|
||||
current_shape = (F, H, W)
|
||||
has_cond = attn_cond is not None
|
||||
if (self.cached_freqs is not None and
|
||||
self.cached_shape == current_shape and
|
||||
self.cached_cond == has_cond):
|
||||
freqs = self.cached_freqs
|
||||
rope_func = "comfy"
|
||||
else:
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
freqs = self.rope_embedder(img_ids).movedim(1, 2)
|
||||
rope_func = "comfy"
|
||||
f_len = ((F + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||
h_len = ((H + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||
w_len = ((W + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||
img_ids = torch.zeros((f_len, h_len, w_len, 3), device=x.device, dtype=x.dtype)
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(0, f_len - 1, steps=f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(0, h_len - 1, steps=h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(0, w_len - 1, steps=w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
|
||||
if attn_cond is not None:
|
||||
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=x.device, dtype=x.dtype)
|
||||
|
||||
#shift
|
||||
shift_f_size = 81 # Default value
|
||||
shift_f = False
|
||||
if shift_f:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(shift_f_size, shift_f_size + cond_f_len - 1,steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
else:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
|
||||
# Combine original and conditional position ids
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
cond_img_ids = repeat(cond_img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
|
||||
|
||||
# Generate RoPE frequencies for the combined positions
|
||||
freqs = self.rope_embedder(combined_img_ids).movedim(1, 2)
|
||||
else:
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
freqs = self.rope_embedder(img_ids).movedim(1, 2)
|
||||
self.cached_freqs = freqs
|
||||
self.cached_shape = current_shape
|
||||
self.cached_cond = has_cond
|
||||
else:
|
||||
rope_func = "default"
|
||||
|
||||
@@ -1494,8 +1504,8 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
token_ref_target_masks = token_ref_target_masks.to(x.dtype).to(device)
|
||||
|
||||
should_calc = True
|
||||
accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device)
|
||||
if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step:
|
||||
accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device)
|
||||
if pred_id is None:
|
||||
pred_id = self.teacache_state.new_prediction(cache_device=self.cache_device)
|
||||
should_calc = True
|
||||
|
||||
Reference in New Issue
Block a user