From 58e5737032c6702a56ff53f7601ba0de8f960356 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 15 Mar 2025 15:45:53 +0200 Subject: [PATCH] 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 --- nodes.py | 24 +++++++---- wanvideo/modules/model.py | 89 +++++++++++++++++++++++++++++++++++---- 2 files changed, 96 insertions(+), 17 deletions(-) diff --git a/nodes.py b/nodes.py index e3f1bbe..a0f720b 100644 --- a/nodes.py +++ b/nodes.py @@ -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) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 52682d6..e56689f 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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: