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:
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user