small possible optimizations

This commit is contained in:
kijai
2025-07-10 14:18:04 +03:00
parent 5df3154112
commit a9ba0e1cb4
+61 -51
View File
@@ -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