add MoCha
This commit is contained in:
@@ -25,6 +25,7 @@ from .cache_methods.nodes_cache import NODE_CLASS_MAPPINGS as NODE_CACHE_CLASS_M
|
||||
from .nodes_deprecated import NODE_CLASS_MAPPINGS as DEPRECATED_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as DEPRECATED_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .s2v.nodes import NODE_CLASS_MAPPINGS as S2V_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as S2V_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .FlashVSR.flashvsr_nodes import NODE_CLASS_MAPPINGS as FLASHVSR_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .mocha.nodes import NODE_CLASS_MAPPINGS as MOCHA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MOCHA_NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
try:
|
||||
from .qwen.qwen import NODE_CLASS_MAPPINGS as QWEN_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as QWEN_NODE_DISPLAY_NAME_MAPPINGS
|
||||
@@ -98,6 +99,7 @@ NODE_CLASS_MAPPINGS.update(SAMPLER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(LYNX_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
@@ -121,5 +123,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(SAMPLER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(LYNX_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(OVI_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
+177
@@ -0,0 +1,177 @@
|
||||
import torch
|
||||
from comfy import model_management as mm
|
||||
import os, gc, math
|
||||
|
||||
def rope_params_mocha(max_seq_len, dim, theta=10000, L_test=25, k=0, start=0):
|
||||
assert dim % 2 == 0
|
||||
exponents = torch.arange(0, dim, 2, dtype=torch.float64).div(dim)
|
||||
inv_theta_pow = 1.0 / torch.pow(theta, exponents)
|
||||
|
||||
if k > 0:
|
||||
print(f"RifleX: Using {k}th freq")
|
||||
inv_theta_pow[k-1] = 0.9 * 2 * torch.pi / L_test
|
||||
|
||||
freqs = torch.outer(torch.arange(start, max_seq_len), inv_theta_pow)
|
||||
freqs = torch.polar(torch.ones_like(freqs), freqs)
|
||||
return freqs
|
||||
|
||||
@torch.autocast(device_type=mm.get_autocast_device(mm.get_torch_device()), enabled=False)
|
||||
@torch.compiler.disable()
|
||||
def rope_apply_mocha(x, grid_sizes, freqs, reverse_time=False):
|
||||
n, c = x.size(2), x.size(3) // 2
|
||||
|
||||
# split freqs
|
||||
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
|
||||
|
||||
# loop over samples
|
||||
output = []
|
||||
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
|
||||
seq_len = f * h * w
|
||||
|
||||
# precompute multipliers
|
||||
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
|
||||
seq_len, n, -1, 2))
|
||||
if reverse_time:
|
||||
time_freqs = freqs[0][:f].view(f, 1, 1, -1)
|
||||
time_freqs = torch.flip(time_freqs, dims=[0])
|
||||
time_freqs = time_freqs.expand(f, h, w, -1)
|
||||
|
||||
spatial_freqs = torch.cat([
|
||||
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
||||
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
||||
], dim=-1)
|
||||
|
||||
freqs_i = torch.cat([time_freqs, spatial_freqs], dim=-1).reshape(seq_len, 1, -1)
|
||||
else:
|
||||
sf = (f - 2) // 2
|
||||
repeat_freqs = torch.cat([
|
||||
freqs[0][1:(1+sf)].view(sf, 1, 1, -1).expand(sf, h, w, -1),
|
||||
freqs[1][1:(1+h)].view(1, h, 1, -1).expand(sf, h, w, -1),
|
||||
freqs[2][1:(1+w)].view(1, 1, w, -1).expand(sf, h, w, -1)
|
||||
], dim=-1)
|
||||
|
||||
mask_freqs = torch.cat([
|
||||
freqs[0][1].view(1, 1, 1, -1).expand(1, h, w, -1),
|
||||
freqs[1][1:(1+h)].view(1, h, 1, -1).expand(1, h, w, -1),
|
||||
freqs[2][1:(1+w)].view(1, 1, w, -1).expand(1, h, w, -1)
|
||||
], dim=-1)
|
||||
|
||||
img_freqs = torch.cat([
|
||||
freqs[0][0].view(1, 1, 1, -1).expand(1, h, w, -1),
|
||||
freqs[1][1:(1+h)].view(1, h, 1, -1).expand(1, h, w, -1),
|
||||
freqs[2][1:(1+w)].view(1, 1, w, -1).expand(1, h, w, -1)
|
||||
], dim=-1)
|
||||
|
||||
if f == 2 * sf + 2:
|
||||
freqs_i = torch.cat([repeat_freqs, repeat_freqs, mask_freqs, img_freqs], dim = 0).reshape(f * h * w, 1, -1).to(x.device)
|
||||
else:
|
||||
bias_freqs = torch.cat([
|
||||
freqs[0][0].view(1, 1, 1, -1).expand(1, h, w, -1),
|
||||
freqs[1][(h+1):(2 * h + 1)].view(1, h, 1, -1).expand(1, h, w, -1),
|
||||
freqs[2][(w+1):(2 * w + 1)].view(1, 1, w, -1).expand(1, h, w, -1)
|
||||
], dim=-1)
|
||||
freqs_i = torch.cat([repeat_freqs, repeat_freqs, mask_freqs, img_freqs, bias_freqs], dim = 0).reshape(f * h * w, 1, -1).to(x.device)
|
||||
|
||||
|
||||
# apply rotary embedding
|
||||
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
|
||||
x_i = torch.cat([x_i, x[i, seq_len:]])
|
||||
|
||||
# append to collection
|
||||
output.append(x_i)
|
||||
return torch.stack(output).to(x.dtype)
|
||||
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
class MochaEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"vae": ("WANVAE",),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
"input_video": ("IMAGE", {"tooltip": "Input video to encode"}),
|
||||
"mask": ("MASK", {"tooltip": "mask"}),
|
||||
"ref1": ("IMAGE", {"tooltip": "Image to encode"}),
|
||||
},
|
||||
"optional": {
|
||||
"ref2": ("IMAGE", {"tooltip": "Image to encode"}),
|
||||
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, vae, force_offload, input_video, mask, ref1, ref2=None, tiled_vae=False):
|
||||
W = input_video.shape[2]
|
||||
H = input_video.shape[1]
|
||||
F = input_video.shape[0]
|
||||
|
||||
lat_h = H // vae.upsampling_factor
|
||||
lat_w = W // vae.upsampling_factor
|
||||
|
||||
F = (F - 1) // 4 * 4 + 1
|
||||
input_video = input_video[: F]
|
||||
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
vae.to(device)
|
||||
|
||||
input_video = input_video.to(device, vae.dtype).unsqueeze(0).permute(0, 4, 1, 2, 3)
|
||||
ref1 = ref1.to(device, vae.dtype).unsqueeze(0).permute(0, 4, 1, 2, 3)
|
||||
if ref2 is not None:
|
||||
ref2 = ref2.to(device, vae.dtype).unsqueeze(0).permute(0, 4, 1, 2, 3)
|
||||
|
||||
|
||||
latents = vae.encode(input_video * 2.0 - 1.0, device, tiled=tiled_vae)
|
||||
|
||||
ref_latents = vae.encode(ref1 * 2.0 - 1.0, device, tiled=tiled_vae)
|
||||
num_refs = 1
|
||||
if ref2 is not None:
|
||||
ref2_latents = vae.encode(ref2 * 2.0 - 1.0, device, tiled=tiled_vae)
|
||||
ref_latents = torch.cat([ref_latents, ref2_latents], dim=2)
|
||||
num_refs = 2
|
||||
|
||||
|
||||
mask = torch.nn.functional.interpolate(mask.unsqueeze(1).to(vae.dtype), size=(lat_h, lat_w), mode='nearest').unsqueeze(1)
|
||||
mask = mask.repeat(1, 16, 1, 1, 1)
|
||||
mask = mask.to(device, vae.dtype)
|
||||
|
||||
mask[mask <= 0.5] = 0
|
||||
mask[mask > 0.5] = 1
|
||||
mask[mask == 0] = -1
|
||||
|
||||
mocha_embeds = torch.cat([latents, mask, ref_latents], dim=2)
|
||||
mocha_embeds = mocha_embeds[0]
|
||||
|
||||
target_shape = (16, (F - 1) // 4 + 1, lat_h, lat_w)
|
||||
|
||||
seq_len = (target_shape[1] * 2 + 1 + num_refs) * (target_shape[2] * target_shape[3] // 4)
|
||||
|
||||
if force_offload:
|
||||
vae.model.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
|
||||
image_embeds = {
|
||||
"seq_len": seq_len,
|
||||
"mocha_embeds": mocha_embeds,
|
||||
"num_frames": F,
|
||||
"target_shape": target_shape,
|
||||
"num_refs": num_refs,
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"MochaEmbeds": MochaEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"MochaEmbeds": "Mocha Embeds",
|
||||
}
|
||||
+25
-4
@@ -283,7 +283,7 @@ class WanVideoSampler:
|
||||
else:
|
||||
cfg = [cfg] * (steps + 1)
|
||||
|
||||
control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = None
|
||||
control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = mocha_embeds = None
|
||||
vace_data = vace_context = vace_scale = None
|
||||
fun_or_fl2v_model = has_ref = drop_last = False
|
||||
phantom_latents = fun_ref_image = ATI_tracks = None
|
||||
@@ -416,6 +416,12 @@ class WanVideoSampler:
|
||||
log.info(f"RecamMaster source video shape: {recam_latents.shape}")
|
||||
seq_len *= 2
|
||||
|
||||
if image_embeds.get("mocha_embeds", None) is not None:
|
||||
mocha_embeds = image_embeds.get("mocha_embeds", None)
|
||||
orig_noise_len = noise.shape[1]
|
||||
seq_len = image_embeds.get("seq_len", seq_len)
|
||||
log.info(f"MoCha embeds shape: {mocha_embeds.shape}")
|
||||
|
||||
# Fun control and control lora
|
||||
control_embeds = image_embeds.get("control_embeds", None)
|
||||
if control_embeds is not None:
|
||||
@@ -1038,6 +1044,18 @@ class WanVideoSampler:
|
||||
transformer.rope_embedder.k = riflex_freq_index
|
||||
transformer.rope_embedder.num_frames = latent_video_length
|
||||
|
||||
if mocha_embeds is not None:
|
||||
from .mocha.nodes import rope_params_mocha
|
||||
log.info(f"Use Mocha RoPE")
|
||||
rope_function = 'mocha'
|
||||
d = transformer.dim // transformer.num_heads
|
||||
freqs = torch.cat([
|
||||
rope_params_mocha(1024, d - 4 * (d // 6), L_test=latent_video_length, k=riflex_freq_index, start=-1),
|
||||
rope_params_mocha(1024, 2 * (d // 6), start=-1),
|
||||
rope_params_mocha(1024, 2 * (d // 6), start=-1)
|
||||
],
|
||||
dim=1)
|
||||
|
||||
transformer.rope_func = rope_function
|
||||
for block in transformer.blocks:
|
||||
block.rope_func = rope_function
|
||||
@@ -1189,6 +1207,9 @@ class WanVideoSampler:
|
||||
if recammaster is not None:
|
||||
z = torch.cat([z, recam_latents.to(z)], dim=1)
|
||||
|
||||
if mocha_embeds is not None:
|
||||
z = torch.cat([z, mocha_embeds.to(z)], dim=1)
|
||||
|
||||
if mtv_input is not None:
|
||||
if ((mtv_start_percent <= current_step_percentage <= mtv_end_percent) or \
|
||||
(mtv_end_percent > 0 and idx == 0 and current_step_percentage >= mtv_start_percent)):
|
||||
@@ -2917,9 +2938,9 @@ class WanVideoSampler:
|
||||
latent = torch.cat(new_latent, dim=1)
|
||||
else:
|
||||
latent = sample_scheduler.step(
|
||||
noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else noise_pred.unsqueeze(0),
|
||||
noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None or mocha_embeds is not None else noise_pred.unsqueeze(0),
|
||||
timestep,
|
||||
latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None else latent.unsqueeze(0),
|
||||
latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None or mocha_embeds is not None else latent.unsqueeze(0),
|
||||
**scheduler_step_args)[0].squeeze(0)
|
||||
if noise_pred_flipped is not None:
|
||||
latent_backwards = sample_scheduler_flipped.step(
|
||||
@@ -2956,7 +2977,7 @@ class WanVideoSampler:
|
||||
current_latent = latent.clone()
|
||||
|
||||
if callback is not None:
|
||||
if recammaster is not None:
|
||||
if recammaster is not None or mocha_embeds is not None:
|
||||
callback_latent = (latent_model_input[:, :orig_noise_len].to(device) - noise_pred[:, :orig_noise_len].to(device) * t.to(device) / 1000).detach()
|
||||
#elif phantom_latents is not None:
|
||||
# callback_latent = (latent_model_input[:,:-phantom_latents.shape[1]].to(device) - noise_pred[:,:-phantom_latents.shape[1]].to(device) * t.to(device) / 1000).detach()
|
||||
|
||||
@@ -1103,6 +1103,10 @@ class WanAttentionBlock(nn.Module):
|
||||
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, reverse_time=reverse_time)
|
||||
k=rope_apply_mocha(k, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
else:
|
||||
q = rope_apply(q, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
k = rope_apply(k, grid_sizes, freqs, reverse_time=reverse_time)
|
||||
|
||||
Reference in New Issue
Block a user