Add option to use comfy native RoPE calculations

comfy rope doesn't use complex numbers and can be torch.compiled, then it's faster and uses less memory
This commit is contained in:
kijai
2025-03-15 15:45:53 +02:00
parent 66893bf39b
commit 58e5737032
2 changed files with 96 additions and 17 deletions
+16 -8
View File
@@ -1242,6 +1242,7 @@ class WanVideoSampler:
"flowedit_args": ("FLOWEDITARGS", ),
"batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Batc cond and uncond for faster sampling, possibly faster on some hardware, uses more memory"}),
"slg_args": ("SLGARGS", ),
"rope_function": (["default", "comfy"], {"default": "default", "tooltip": "!EXPERIMENTAL! Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}),
}
}
@@ -1252,7 +1253,7 @@ class WanVideoSampler:
def process(self, model, text_embeds, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index,
force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None,
teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None):
teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default"):
#assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache."
patcher = model
model = model.model
@@ -1426,13 +1427,20 @@ class WanVideoSampler:
latent = noise.to(device)
d = transformer.dim // transformer.num_heads
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
],
dim=1)
freqs = None
transformer.rope_embedder.k = None
transformer.rope_embedder.num_frames = None
if rope_function=="comfy":
transformer.rope_embedder.k = riflex_freq_index
transformer.rope_embedder.num_frames = latent_video_length
else:
d = transformer.dim // transformer.num_heads
freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
],
dim=1)
if not isinstance(cfg, list):
cfg = [cfg] * (steps +1)
+80 -9
View File
@@ -5,7 +5,7 @@ import torch
import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.modeling_utils import ModelMixin
from einops import repeat
from ...enhance_a_video.enhance import get_feta_scores
from ...enhance_a_video.globals import is_enhance_enabled
@@ -18,6 +18,45 @@ import gc
import comfy.model_management as mm
from ...utils import log, get_module_memory_mb
from comfy.ldm.flux.math import apply_rope as apply_rope_comfy
def rope_riflex(pos, dim, theta, L_test, k):
from einops import rearrange
assert dim % 2 == 0
if mm.is_device_mps(pos.device) or mm.is_intel_xpu() or mm.is_directml_enabled():
device = torch.device("cpu")
else:
device = pos.device
scale = torch.linspace(0, (dim - 2) / dim, steps=dim//2, dtype=torch.float64, device=device)
omega = 1.0 / (theta**scale)
# RIFLEX modification - adjust last frequency component if L_test and k are provided
if k and L_test:
omega[k-1] = 0.9 * 2 * torch.pi / L_test
out = torch.einsum("...n,d->...nd", pos.to(dtype=torch.float32, device=device), omega)
out = torch.stack([torch.cos(out), -torch.sin(out), torch.sin(out), torch.cos(out)], dim=-1)
out = rearrange(out, "b n d (i j) -> b n d i j", i=2, j=2)
return out.to(dtype=torch.float32, device=pos.device)
class EmbedND_RifleX(nn.Module):
def __init__(self, dim, theta, axes_dim, num_frames, k):
super().__init__()
self.dim = dim
self.theta = theta
self.axes_dim = axes_dim
self.num_frames = num_frames
self.k = k
def forward(self, ids):
n_axes = ids.shape[-1]
emb = torch.cat(
[rope_riflex(ids[..., i], self.axes_dim[i], self.theta, self.num_frames, self.k if i == 0 else 0) for i in range(n_axes)],
dim=-3,
)
return emb.unsqueeze(1)
def poly1d(coefficients, x):
result = torch.zeros_like(x)
for i, coeff in enumerate(coefficients):
@@ -182,7 +221,7 @@ class WanSelfAttention(nn.Module):
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def forward(self, x, seq_lens, grid_sizes, freqs):
def forward(self, x, seq_lens, grid_sizes, freqs, rope_func = "default"):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -222,8 +261,11 @@ class WanSelfAttention(nn.Module):
).permute(0, 2, 1, 3)
#print("inner attention", x.shape) #inner attention torch.Size([1, 12, 32760, 128])
else:
q=rope_apply(q, grid_sizes, freqs)
k=rope_apply(k, grid_sizes, freqs)
if rope_func == "comfy":
q, k = apply_rope_comfy(q, k, freqs)
else:
q=rope_apply(q, grid_sizes, freqs)
k=rope_apply(k, grid_sizes, freqs)
if is_enhance_enabled():
feta_scores = get_feta_scores(q, k)
@@ -374,6 +416,7 @@ class WanAttentionBlock(nn.Module):
freqs,
context,
context_lens,
rope_func = "default",
):
r"""
Args:
@@ -390,7 +433,7 @@ class WanAttentionBlock(nn.Module):
# self-attention
y = self.self_attn(
self.norm1(x).float() * (1 + e[1]) + e[0], seq_lens, grid_sizes,
freqs)
freqs, rope_func=rope_func)
x = x.to(torch.float32) + (y.to(torch.float32) * e[2].to(torch.float32))
# cross-attention & ffn function
@@ -585,6 +628,15 @@ class WanModel(ModelMixin, ConfigMixin):
# head
self.head = Head(dim, out_dim, patch_size, eps)
d = self.dim // self.num_heads
self.rope_embedder = EmbedND_RifleX(
d,
10000.0,
[d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)],
num_frames=None,
k=None,
)
# buffers (don't use register_buffer otherwise dtype will be changed in to())
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
@@ -665,9 +717,11 @@ class WanModel(ModelMixin, ConfigMixin):
if self.model_type == 'i2v':
assert clip_fea is not None and y is not None
# params
#device = self.patch_embedding.weight.device
if freqs.device != device:
freqs = freqs.to(device)
device = self.patch_embedding.weight.device
if freqs is not None and freqs.device != device:
freqs = freqs.to(device)
_, F, H, W = x[0].shape
if y is not None:
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
@@ -696,6 +750,21 @@ class WanModel(ModelMixin, ConfigMixin):
dim=1) for u in x
])
if freqs is None: #comfy rope
rope_func = "comfy"
f_len = ((F + (self.patch_size[0] // 2)) // self.patch_size[0])
h_len = ((H + (self.patch_size[1] // 2)) // self.patch_size[1])
w_len = ((W + (self.patch_size[2] // 2)) // self.patch_size[2])
img_ids = torch.zeros((f_len, h_len, w_len, 3), device=x.device, dtype=x.dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(0, f_len - 1, steps=f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(0, h_len - 1, steps=h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(0, w_len - 1, steps=w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
freqs = self.rope_embedder(img_ids).movedim(1, 2)
else:
rope_func = "default"
# time embeddings
with torch.autocast(device_type='cuda', dtype=torch.float32):
e = self.time_embedding(
@@ -773,7 +842,9 @@ class WanModel(ModelMixin, ConfigMixin):
grid_sizes=grid_sizes,
freqs=freqs,
context=context,
context_lens=context_lens)
context_lens=context_lens,
rope_func=rope_func
)
for b, block in enumerate(self.blocks):
if self.slg_blocks is not None: