Reduce needless torch.compile recompiles

This commit is contained in:
kijai
2025-10-22 13:29:13 +03:00
parent aa610cde2b
commit 67fcf0ba52
4 changed files with 28 additions and 21 deletions
+3 -2
View File
@@ -1,7 +1,6 @@
import torch
import torch.nn as nn
from accelerate import init_empty_weights
from comfy.ops import cast_bias_weight
#based on https://github.com/huggingface/diffusers/blob/main/src/diffusers/quantizers/gguf/utils.py
def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, scale_weights=None):
@@ -85,7 +84,9 @@ class CustomLinear(nn.Linear):
self.weight_function = []
def forward(self, input):
weight, bias = cast_bias_weight(self, input)
if self.bias is not None:
bias = self.bias.to(input)
weight = self.weight.to(input)
if self.scale_weight is not None:
if weight.numel() < input.numel():
+1 -1
View File
@@ -328,7 +328,7 @@ class WanVideoTorchCompileSettings:
},
"optional": {
"dynamo_recompile_limit": ("INT", {"default": 128, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.recompile_limit"}),
"force_parameter_static_shapes": ("BOOLEAN", {"default": True, "tooltip": "torch._dynamo.config.force_parameter_static_shapes"}),
"force_parameter_static_shapes": ("BOOLEAN", {"default": False, "tooltip": "torch._dynamo.config.force_parameter_static_shapes"}),
},
}
RETURN_TYPES = ("WANCOMPILEARGS",)
+1
View File
@@ -505,6 +505,7 @@ def compile_model(transformer, compile_args=None):
if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'):
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
torch._dynamo.config.force_parameter_static_shapes = compile_args["force_parameter_static_shapes"]
torch._dynamo.config.allow_unspec_int_on_nn_module = True
try:
torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"]
except Exception as e:
+23 -18
View File
@@ -1025,6 +1025,7 @@ class WanAttentionBlock(nn.Module):
x_ip=None, e_ip=None, freqs_ip=None, ip_scale=1.0, #stand-in
adapter_proj=None, #fantasyportrait
reverse_time=False,
zero_timestep=False, #s2v zero timestep
mtv_motion_tokens=None, mtv_motion_rotary_emb=None, mtv_strength=1.0, mtv_freqs=None, #mtv crafter
humo_audio_input=None, humo_audio_scale=1.0, #humo audio
lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx
@@ -1038,9 +1039,8 @@ class WanAttentionBlock(nn.Module):
grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W)
freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2]
"""
self.original_seq_len = original_seq_len
self.zero_timestep = len(e) == 2
if self.zero_timestep: #s2v zero timestep
zero_timestep = len(e) == 2
if zero_timestep: #s2v zero timestep
self.seg_idx = e[1]
self.seg_idx = min(max(0, self.seg_idx), x.size(1))
self.seg_idx = [0, self.seg_idx, x.size(1)]
@@ -1186,7 +1186,7 @@ class WanAttentionBlock(nn.Module):
)
# S2V
if self.zero_timestep:
if zero_timestep:
z = []
for i in range(2):
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_msa[:, i:i + 1])
@@ -1229,9 +1229,9 @@ class WanAttentionBlock(nn.Module):
x = self.cross_attn_ffn(x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
audio_proj, audio_scale, num_latent_frames, nag_params, nag_context, is_uncond,
multitalk_audio_embedding, x_ref_attn_map, human_num, inner_t, inner_c, cross_freqs,
adapter_proj=adapter_proj, ip_scale=ip_scale,
adapter_proj=adapter_proj, ip_scale=ip_scale, zero_timestep=zero_timestep,
mtv_freqs=mtv_freqs, mtv_motion_tokens=mtv_motion_tokens, mtv_motion_rotary_emb=mtv_motion_rotary_emb, mtv_strength=mtv_strength,
humo_audio_input=humo_audio_input, humo_audio_scale=humo_audio_scale, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale
humo_audio_input=humo_audio_input, humo_audio_scale=humo_audio_scale, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, original_seq_len=original_seq_len
)
else:
if self.rope_func == "comfy_chunked":
@@ -1251,15 +1251,15 @@ class WanAttentionBlock(nn.Module):
def cross_attn_ffn(self, x, context, grid_sizes, shift_mlp, scale_mlp, gate_mlp, clip_embed,
audio_proj, audio_scale, num_latent_frames, nag_params,
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num,
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale, mtv_freqs, mtv_motion_tokens, mtv_motion_rotary_emb, mtv_strength,
humo_audio_input, humo_audio_scale, lynx_x_ip, lynx_ip_scale):
nag_context, is_uncond, multitalk_audio_embedding, x_ref_attn_map, human_num,
inner_t, inner_c, cross_freqs, adapter_proj, ip_scale, zero_timestep, mtv_freqs, mtv_motion_tokens, mtv_motion_rotary_emb, mtv_strength,
humo_audio_input, humo_audio_scale, lynx_x_ip, lynx_ip_scale, original_seq_len):
x = x + self.cross_attn(self.norm3(x), 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, is_uncond=is_uncond,
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=self.original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, )
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, )
# MultiTalk
if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock):
x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding,
@@ -1275,11 +1275,11 @@ class WanAttentionBlock(nn.Module):
if humo_audio_input is not None:
x = self.audio_cross_attn_wrapper(x, humo_audio_input, grid_sizes, humo_audio_scale)
if self.rope_func == "comfy_chunked" and not self.zero_timestep:
if self.rope_func == "comfy_chunked" and not zero_timestep:
y = self.ffn_chunked(x, shift_mlp, scale_mlp)
else:
norm2_x = self.norm2(x)
if self.zero_timestep:
if zero_timestep:
parts = []
for i in range(2):
parts.append(norm2_x[:, self.seg_idx[i]:self.seg_idx[i + 1]] *
@@ -1290,7 +1290,7 @@ class WanAttentionBlock(nn.Module):
input_x = torch.addcmul(shift_mlp, norm2_x, 1 + scale_mlp)
del shift_mlp, scale_mlp, norm2_x
y = self.ffn(input_x)
if self.zero_timestep:
if zero_timestep:
z = []
for i in range(2):
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_mlp[:, i:i + 1])
@@ -1371,8 +1371,10 @@ class VaceWanAttentionBlock(WanAttentionBlock):
rms_norm_function="default"
):
super().__init__(cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm, cross_attn_norm, eps, attention_mode, rope_func, rms_norm_function=rms_norm_function)
self.block_id = block_id
if block_id == 0:
self.register_buffer('block_id', torch.tensor(block_id, dtype=torch.long))
if torch.equal(self.block_id, torch.tensor(0)):
self.before_proj = nn.Linear(in_features, out_features)
self.after_proj = nn.Linear(in_features, out_features)
@@ -1402,7 +1404,10 @@ class BaseWanAttentionBlock(WanAttentionBlock):
super().__init__(cross_attn_type, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm,
cross_attn_norm, eps, attention_mode, rope_func, rms_norm_function=rms_norm_function,
block_idx=block_idx, lynx_ip_layers=lynx_ip_layers, lynx_ref_layers=lynx_ref_layers)
self.block_id = block_id
if block_id is not None:
self.register_buffer('block_id', torch.tensor(block_id, dtype=torch.long))
else:
self.block_id = None
def forward(self, x, vace_hints=None, vace_context_scale=[1.0], **kwargs):
x, x_ip, lynx_ref_feature, x_ovi = super().forward(x, **kwargs)
@@ -2721,8 +2726,8 @@ class WanModel(torch.nn.Module):
freqs=freqs,
context=context,
clip_embed=clip_embed,
current_step=current_step,
last_step=last_step,
current_step=torch.tensor(current_step),
last_step=torch.tensor(last_step, dtype=torch.bool),
video_attention_split_steps=self.video_attention_split_steps,
camera_embed=camera_embed,
audio_proj=audio_proj,