Reduce needless torch.compile recompiles
This commit is contained in:
+3
-2
@@ -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():
|
||||
|
||||
@@ -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",)
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user