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.
This commit is contained in:
kijai
2025-11-30 17:14:53 +02:00
parent 772642b4f1
commit 99c3978da4
+94 -121
View File
@@ -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: