From 67fcf0ba529b7e646eb896371fdb3c7fb48193a0 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 22 Oct 2025 13:29:13 +0300 Subject: [PATCH] Reduce needless torch.compile recompiles --- custom_linear.py | 5 +++-- nodes_model_loading.py | 2 +- utils.py | 1 + wanvideo/modules/model.py | 41 ++++++++++++++++++++++----------------- 4 files changed, 28 insertions(+), 21 deletions(-) diff --git a/custom_linear.py b/custom_linear.py index 81c498a..7154b2a 100644 --- a/custom_linear.py +++ b/custom_linear.py @@ -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(): diff --git a/nodes_model_loading.py b/nodes_model_loading.py index fc14f1b..00eec51 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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",) diff --git a/utils.py b/utils.py index cedf537..a21302c 100644 --- a/utils.py +++ b/utils.py @@ -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: diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 7f8a196..9264a73 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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,