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:
+92
-119
@@ -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:
|
||||
@@ -1199,8 +1171,7 @@ class WanAttentionBlock(nn.Module):
|
||||
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:
|
||||
|
||||
Reference in New Issue
Block a user