From 99c3978da4a55a03249669bef5647d7dbda7a5d1 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 30 Nov 2025 17:14:53 +0200 Subject: [PATCH 1/4] Reduce peak VRAM usage when not using torch.compile (and some even with it) Found some intermediates that weren't freed which should reduce VRAM usage overall, and modified RoPE application outside torch compile for similar gains than when using torch.compile. --- wanvideo/modules/model.py | 215 +++++++++++++++++--------------------- 1 file changed, 94 insertions(+), 121 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index cdb2341..87c2e21 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -23,7 +23,8 @@ from ...multitalk.multitalk import get_attn_map_with_target from ...echoshot.echoshot import rope_apply_z, rope_apply_c, rope_apply_echoshot from ...MTV.mtv import apply_rotary_emb - +from comfy.ldm.flux.math import apply_rope1 as apply_rope_comfy1 +from comfy.ldm.flux.math import apply_rope as apply_rope_comfy from comfy import model_management as mm __all__ = ['WanModel'] @@ -120,70 +121,6 @@ def torch_dfs(model: nn.Module, parent_name='root'): modules += child_modules return modules, module_names -#from comfy.ldm.flux.math import apply_rope as apply_rope_comfy -def apply_rope_comfy(xq, xk, freqs_cis): - xq_ = xq.to(dtype=freqs_cis.dtype).reshape(*xq.shape[:-1], -1, 1, 2) - xk_ = xk.to(dtype=freqs_cis.dtype).reshape(*xk.shape[:-1], -1, 1, 2) - xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1] - xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1] - return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk) - -def apply_rope_comfy_chunked(xq, xk, freqs_cis, num_chunks=4): - seq_dim = 1 - - # Initialize output tensors - xq_out = torch.empty_like(xq) - xk_out = torch.empty_like(xk) - - # Calculate chunks - seq_len = xq.shape[seq_dim] - chunk_sizes = [seq_len // num_chunks + (1 if i < seq_len % num_chunks else 0) - for i in range(num_chunks)] - - # First pass: process xq completely - start_idx = 0 - for size in chunk_sizes: - end_idx = start_idx + size - - slices = [slice(None)] * len(xq.shape) - slices[seq_dim] = slice(start_idx, end_idx) - - freq_slices = [slice(None)] * len(freqs_cis.shape) - if seq_dim < len(freqs_cis.shape): - freq_slices[seq_dim] = slice(start_idx, end_idx) - freqs_chunk = freqs_cis[tuple(freq_slices)] - - xq_chunk = xq[tuple(slices)] - xq_chunk_ = xq_chunk.to(dtype=freqs_cis.dtype).reshape(*xq_chunk.shape[:-1], -1, 1, 2) - xq_out[tuple(slices)] = (freqs_chunk[..., 0] * xq_chunk_[..., 0] + - freqs_chunk[..., 1] * xq_chunk_[..., 1]).reshape(*xq_chunk.shape).type_as(xq) - - del xq_chunk, xq_chunk_, freqs_chunk - start_idx = end_idx - - # Second pass: process xk completely - start_idx = 0 - for size in chunk_sizes: - end_idx = start_idx + size - - slices = [slice(None)] * len(xk.shape) - slices[seq_dim] = slice(start_idx, end_idx) - - freq_slices = [slice(None)] * len(freqs_cis.shape) - if seq_dim < len(freqs_cis.shape): - freq_slices[seq_dim] = slice(start_idx, end_idx) - freqs_chunk = freqs_cis[tuple(freq_slices)] - - xk_chunk = xk[tuple(slices)] - xk_chunk_ = xk_chunk.to(dtype=freqs_cis.dtype).reshape(*xk_chunk.shape[:-1], -1, 1, 2) - xk_out[tuple(slices)] = (freqs_chunk[..., 0] * xk_chunk_[..., 0] + - freqs_chunk[..., 1] * xk_chunk_[..., 1]).reshape(*xk_chunk.shape).type_as(xk) - - del xk_chunk, xk_chunk_, freqs_chunk - start_idx = end_idx - - return xq_out, xk_out - def rope_riflex(pos, dim, i, theta, L_test, k, ntk_factor=1.0): assert dim % 2 == 0 if mm.is_device_mps(pos.device) or mm.is_intel_xpu() or mm.is_directml_enabled(): @@ -462,21 +399,59 @@ class WanSelfAttention(nn.Module): self.norm_q = WanRMSNorm(norm_dim, eps=eps) if qk_norm else nn.Identity() self.norm_k = WanRMSNorm(norm_dim, eps=eps) if qk_norm else nn.Identity() - def qkv_fn(self, x): + def qkv_fn(self, x, is_longcat=False): b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim - q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d) - k = self.norm_k(self.k(x).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d) + if is_longcat: + q = self.q(x).view(b, s, n, d) + q = self.norm_q(q.float()).to(x.dtype) + k = self.k(x).view(b, s, n, d) + k = self.norm_k(k.float()).to(x.dtype) + else: + q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d) + k = self.norm_k(self.k(x).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, s, n, d) v = self.v(x).view(b, s, n, d) return q, k, v - - def qkv_fn_longcat(self, x): + + def qkv_fn_qk_with_rope(self, x, layer, freqs, num_chunks=1, is_longcat=False): b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim - q = self.q(x).view(b, s, n, d) - q = self.norm_q(q.float()).to(x.dtype) - k = self.k(x).view(b, s, n, d) - k = self.norm_k(k.float()).to(x.dtype) - v = self.v(x).view(b, s, n, d) - return q, k, v + + linear_layer = self.q if layer == 'q' else self.k + norm_layer = self.norm_q if layer == 'q' else self.norm_k + + use_chunked = num_chunks > 1 + if use_chunked: + chunk_sizes = [s // num_chunks + (1 if i < s % num_chunks else 0) + for i in range(num_chunks)] + + out = torch.empty(b, s, n, d, dtype=x.dtype, device=x.device) + start_idx = 0 + for size in chunk_sizes: + end_idx = start_idx + size + + x_chunk = x[:, start_idx:end_idx] + + if is_longcat: + chunk = linear_layer(x_chunk).view(b, size, n, d) + chunk = norm_layer(chunk.float()).to(x.dtype) + else: + chunk = norm_layer(linear_layer(x_chunk).to(norm_layer.weight.dtype)).to(x.dtype).view(b, size, n, d) + + freqs_chunk = freqs[:, start_idx:end_idx] if freqs.shape[1] > 1 else freqs + out[:, start_idx:end_idx] = apply_rope_comfy1(chunk, freqs_chunk) + + start_idx = end_idx + return out + else: + if is_longcat: + result = linear_layer(x).view(b, s, n, d) + result = norm_layer(result.float()).to(x.dtype) + else: + result = norm_layer(linear_layer(x).to(norm_layer.weight.dtype)).to(x.dtype).view(b, s, n, d) + return apply_rope_comfy1(result, freqs) + + def qkv_fn_v(self, x): + b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim + return self.v(x).view(b, s, n, d) def qkv_fn_ip(self, x): b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim @@ -935,9 +910,7 @@ class WanAttentionBlock(nn.Module): self.norm2 = WanLayerNorm(self.dim, eps) if not is_longcat: - self.ffn = nn.Sequential( - nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), - nn.Linear(ffn2_dim, out_features)) + self.ffn = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features)) else: from ...LongCat.layers import FeedForwardSwiGLU mlp_ratio = 4 @@ -1002,23 +975,17 @@ class WanAttentionBlock(nn.Module): else: return torch.addcmul(shift_msa, norm_x, 1 + scale_msa) - def ffn_chunked(self, x, shift_mlp, scale_mlp, num_chunks=4): - modulated_input = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp).to(x.dtype) + def ffn_chunked(self, mod_x, num_chunks=4): + seq_len = mod_x.shape[1] + if seq_len <= 8192 or num_chunks <= 1: + return self.ffn(mod_x) - result = torch.empty_like(x) - seq_len = modulated_input.shape[1] + chunk_size = (seq_len + num_chunks - 1) // num_chunks + for i in range(0, seq_len, chunk_size): + end_idx = min(i + chunk_size, seq_len) + mod_x[:, i:end_idx] = self.ffn(mod_x[:, i:end_idx].contiguous()) - chunk_sizes = [seq_len // num_chunks + (1 if i < seq_len % num_chunks else 0) - for i in range(num_chunks)] - - start_idx = 0 - for size in chunk_sizes: - end_idx = start_idx + size - chunk = modulated_input[:, start_idx:end_idx, :] - result[:, start_idx:end_idx, :] = self.ffn(chunk) - start_idx = end_idx - - return result + return mod_x #region attention forward def forward( @@ -1099,6 +1066,9 @@ class WanAttentionBlock(nn.Module): # self-attention variables q_ip = k_ip = v_ip = None + if lynx_ref_feature is None and self.self_attn.ref_adapter is not None: + lynx_ref_feature = input_x + #RoPE and QKV computation if inner_t is not None: #query, key, value @@ -1109,33 +1079,35 @@ class WanAttentionBlock(nn.Module): # First pass - separate main and IP components x_main, x_ip_input = input_x[:, : -self.cond_size], input_x[:, -self.cond_size :] # Compute QKV for main content - q, k, v = self.self_attn.qkv_fn(x_main) if self.rope_func == "comfy": - q, k = apply_rope_comfy(q, k, freqs) + q = self.self_attn.qkv_fn_qk_with_rope(x_main, "q", freqs) + k = self.self_attn.qkv_fn_qk_with_rope(x_main, "k", freqs) + v = self.self_attn.qkv_fn_v(x_main) elif self.rope_func == "comfy_chunked": - q, k = apply_rope_comfy_chunked(q, k, freqs) + q = self.self_attn.qkv_fn_qk_with_rope(x_main, "q", freqs, num_chunks=2) + k = self.self_attn.qkv_fn_qk_with_rope(x_main, "k", freqs, num_chunks=2) + v = self.self_attn.qkv_fn_v(x_main) # Compute QKV for IP content - q_ip, k_ip, v_ip = self.self_attn.qkv_fn_ip(x_ip_input) - if self.rope_func == "comfy": + if "comfy" in self.rope_func: + q_ip, k_ip, v_ip = self.self_attn.qkv_fn_ip(x_ip_input) q_ip, k_ip = apply_rope_comfy(q_ip, k_ip, freqs_ip) - elif self.rope_func == "comfy_chunked": - q_ip, k_ip = apply_rope_comfy_chunked(q_ip, k_ip, freqs_ip) else: - if is_longcat: - q, k, v = self.self_attn.qkv_fn_longcat(input_x) + if "comfy" in self.rope_func: + num_chunks = 2 if self.rope_func == "comfy_chunked" else 1 + q = self.self_attn.qkv_fn_qk_with_rope(input_x, "q", freqs, num_chunks=num_chunks, is_longcat=is_longcat) + k = self.self_attn.qkv_fn_qk_with_rope(input_x, "k", freqs, num_chunks=num_chunks, is_longcat=is_longcat) + v = self.self_attn.qkv_fn_v(input_x) else: q, k, v = self.self_attn.qkv_fn(input_x) - if self.rope_func == "comfy": - q, k = apply_rope_comfy(q, k, freqs) - elif self.rope_func == "comfy_chunked": - q, k = apply_rope_comfy_chunked(q, k, freqs) - elif self.rope_func == "mocha": - from ...mocha.nodes import rope_apply_mocha - q=rope_apply_mocha(q, grid_sizes, freqs) - k=rope_apply_mocha(k, grid_sizes, freqs) - else: - q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time) - k = rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time) + if self.rope_func == "mocha": + from ...mocha.nodes import rope_apply_mocha + q = rope_apply_mocha(q, grid_sizes, freqs) + k = rope_apply_mocha(k, grid_sizes, freqs) + else: + q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time) + k = rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time) + + del input_x if x_ovi is not None: q_ovi, k_ovi, v_ovi = self.audio_block.self_attn.qkv_fn(input_x_ovi) @@ -1143,7 +1115,7 @@ class WanAttentionBlock(nn.Module): k_ovi = rope_apply(k_ovi, grid_sizes_ovi, freqs_ovi) y_ovi = self.audio_block.self_attn.forward(q_ovi, k_ovi, v_ovi, seq_lens_ovi) x_ovi = x_ovi.addcmul(y_ovi, gate_msa_ovi) - + del input_x_ovi, y_ovi, gate_msa_ovi # FETA if enhance_enabled: @@ -1198,9 +1170,8 @@ class WanAttentionBlock(nn.Module): y = torch.cat([x_cond, x_noise], dim=1).contiguous() else: y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale) - - if lynx_ref_feature is None and self.self_attn.ref_adapter is not None: - lynx_ref_feature = input_x + + del q, k, v # FETA if enhance_enabled: @@ -1281,7 +1252,8 @@ class WanAttentionBlock(nn.Module): # ffn if self.rope_func == "comfy_chunked": - x_ffn = self.ffn_chunked(x, shift_mlp, scale_mlp) + mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp) + x_ffn = self.ffn_chunked(mod_x) else: if zero_timestep: norm2_x = self.norm2(x) @@ -1296,8 +1268,9 @@ class WanAttentionBlock(nn.Module): mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp) else: mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C) - x_ffn = self.ffn(mod_x.to(input_dtype)) - del shift_mlp, scale_mlp + del shift_mlp, scale_mlp + x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1) + del mod_x # gate_mlp if zero_timestep: From 1e9e2be62281de76b9a761082d31c3654cd2b381 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 30 Nov 2025 17:32:24 +0200 Subject: [PATCH 2/4] Avoid recompile here --- wanvideo/modules/model.py | 23 +++++++++++++---------- 1 file changed, 13 insertions(+), 10 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 87c2e21..1ae5100 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -412,11 +412,8 @@ class WanSelfAttention(nn.Module): v = self.v(x).view(b, s, n, d) return q, k, v - def qkv_fn_qk_with_rope(self, x, layer, freqs, num_chunks=1, is_longcat=False): + def _qkv_fn_with_rope(self, x, linear_layer, norm_layer, freqs, num_chunks=1, is_longcat=False): b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim - - linear_layer = self.q if layer == 'q' else self.k - norm_layer = self.norm_q if layer == 'q' else self.norm_k use_chunked = num_chunks > 1 if use_chunked: @@ -449,6 +446,12 @@ class WanSelfAttention(nn.Module): result = norm_layer(linear_layer(x).to(norm_layer.weight.dtype)).to(x.dtype).view(b, s, n, d) return apply_rope_comfy1(result, freqs) + def qkv_fn_q_with_rope(self, x, freqs, num_chunks=1, is_longcat=False): + return self._qkv_fn_with_rope(x, self.q, self.norm_q, freqs, num_chunks, is_longcat) + + def qkv_fn_k_with_rope(self, x, freqs, num_chunks=1, is_longcat=False): + return self._qkv_fn_with_rope(x, self.k, self.norm_k, freqs, num_chunks, is_longcat) + def qkv_fn_v(self, x): b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim return self.v(x).view(b, s, n, d) @@ -1080,12 +1083,12 @@ class WanAttentionBlock(nn.Module): x_main, x_ip_input = input_x[:, : -self.cond_size], input_x[:, -self.cond_size :] # Compute QKV for main content if self.rope_func == "comfy": - q = self.self_attn.qkv_fn_qk_with_rope(x_main, "q", freqs) - k = self.self_attn.qkv_fn_qk_with_rope(x_main, "k", freqs) + q = self.self_attn.qkv_fn_q_with_rope(x_main, freqs) + k = self.self_attn.qkv_fn_k_with_rope(x_main, freqs) v = self.self_attn.qkv_fn_v(x_main) elif self.rope_func == "comfy_chunked": - q = self.self_attn.qkv_fn_qk_with_rope(x_main, "q", freqs, num_chunks=2) - k = self.self_attn.qkv_fn_qk_with_rope(x_main, "k", freqs, num_chunks=2) + q = self.self_attn.qkv_fn_q_with_rope(x_main, freqs, num_chunks=2) + k = self.self_attn.qkv_fn_k_with_rope(x_main, freqs, num_chunks=2) v = self.self_attn.qkv_fn_v(x_main) # Compute QKV for IP content if "comfy" in self.rope_func: @@ -1094,8 +1097,8 @@ class WanAttentionBlock(nn.Module): else: if "comfy" in self.rope_func: num_chunks = 2 if self.rope_func == "comfy_chunked" else 1 - q = self.self_attn.qkv_fn_qk_with_rope(input_x, "q", freqs, num_chunks=num_chunks, is_longcat=is_longcat) - k = self.self_attn.qkv_fn_qk_with_rope(input_x, "k", freqs, num_chunks=num_chunks, is_longcat=is_longcat) + q = self.self_attn.qkv_fn_q_with_rope(input_x, freqs, num_chunks=num_chunks, is_longcat=is_longcat) + k = self.self_attn.qkv_fn_k_with_rope(input_x, freqs, num_chunks=num_chunks, is_longcat=is_longcat) v = self.self_attn.qkv_fn_v(input_x) else: q, k, v = self.self_attn.qkv_fn(input_x) From a9cd073f299eab885da400ed2667bddba4629f63 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 30 Nov 2025 17:52:56 +0200 Subject: [PATCH 3/4] Remove unnecessary recompile when using cfg --- wanvideo/modules/model.py | 19 +++++++++---------- 1 file changed, 9 insertions(+), 10 deletions(-) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 1ae5100..9c26fe5 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -640,7 +640,7 @@ class WanT2VCrossAttention(WanSelfAttention): self.k_fusion = None def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0, - num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy", + num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy", inner_t=None, inner_c=None, cross_freqs=None, adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, num_cond_latents=None, **kwargs): b, n, d = x.size(0), self.num_heads, self.head_dim @@ -656,7 +656,7 @@ class WanT2VCrossAttention(WanSelfAttention): else: q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d) - if nag_context is not None and not is_uncond: + if nag_context is not None: x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) else: if is_longcat: @@ -667,7 +667,7 @@ class WanT2VCrossAttention(WanSelfAttention): v = self.v(context).view(b, -1, n, d) #EchoShot rope - if inner_t is not None and cross_freqs is not None and not is_uncond: + if inner_t is not None and cross_freqs is not None: q = rope_apply_z(q, grid_sizes, cross_freqs, inner_t).to(q) k = rope_apply_c(k, cross_freqs, inner_c).to(q) @@ -736,7 +736,7 @@ class WanI2VCrossAttention(WanSelfAttention): self.attention_mode = attention_mode def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, - audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy", + audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, rope_func="comfy", adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs): r""" Args: @@ -747,7 +747,7 @@ class WanI2VCrossAttention(WanSelfAttention): # compute query q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype) - if nag_context is not None and not is_uncond: + if nag_context is not None: x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) else: # text attention @@ -1003,7 +1003,6 @@ class WanAttentionBlock(nn.Module): original_seq_len=None, enhance_enabled=False, #feta nag_params={}, nag_context=None, #normalized attention guidance - is_uncond=False, multitalk_audio_embedding=None, ref_target_masks=None, human_num=0, #multitalk inner_t=None, inner_c=None, cross_freqs=None, #echoshot x_ip=None, e_ip=None, freqs_ip=None, ip_scale=1.0, #stand-in @@ -1233,7 +1232,7 @@ class WanAttentionBlock(nn.Module): return x, x_ip, lynx_ref_feature, x_ovi else: x = x + self.cross_attn(self.norm3(x.to(self.norm3.weight.dtype)).to(input_dtype), 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, + num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, 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=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents) x = x.to(input_dtype) @@ -2800,13 +2799,13 @@ class WanModel(torch.nn.Module): original_seq_len=self.original_seq_len, enhance_enabled=enhance_enabled, audio_scale=audio_scale, - nag_params=nag_params, nag_context=nag_context, - is_uncond = is_uncond, + nag_params=nag_params, + nag_context=nag_context if not is_uncond else None, multitalk_audio_embedding=multitalk_audio_embedding if multitalk_audio is not None else None, ref_target_masks=token_ref_target_masks if multitalk_audio is not None else None, human_num=human_num if multitalk_audio is not None else 0, inner_t=inner_t, inner_c=inner_c, - cross_freqs=self.cross_freqs if inner_t is not None else None, + cross_freqs=self.cross_freqs if inner_t is not None and not is_uncond else None, freqs_ip=freqs_ip if x_ip is not None else None, e_ip=e0_ip if x_ip is not None else None, adapter_proj=adapter_proj, From 0cba1edd4eeab8bc8b4d0295cfbcf3f33909ec53 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 30 Nov 2025 17:53:28 +0200 Subject: [PATCH 4/4] Better just not compile this as it's causing issues --- custom_linear.py | 1 + 1 file changed, 1 insertion(+) diff --git a/custom_linear.py b/custom_linear.py index e6cd2bf..5b047b1 100644 --- a/custom_linear.py +++ b/custom_linear.py @@ -96,6 +96,7 @@ class CustomLinear(nn.Linear): if not allow_compile: self._get_weight_with_lora = torch.compiler.disable()(self._get_weight_with_lora) + self.forward = torch.compiler.disable()(self.forward) def set_lora_diffs(self, lora_diffs, device=torch.device("cpu")): self.lora_diffs = []