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] 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: