From 88defbfdd1ba80e9aaf5cee86d6ed5d10bc6f8e5 Mon Sep 17 00:00:00 2001 From: unrealMJ Date: Tue, 21 Oct 2025 09:40:54 +0800 Subject: [PATCH] add MoCha --- __init__.py | 3 + mocha/nodes.py | 177 ++++++++++++++++++++++++++++++++++++++ nodes_sampler.py | 29 ++++++- wanvideo/modules/model.py | 4 + 4 files changed, 209 insertions(+), 4 deletions(-) create mode 100644 mocha/nodes.py diff --git a/__init__.py b/__init__.py index d12d003..41f629c 100644 --- a/__init__.py +++ b/__init__.py @@ -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"] \ No newline at end of file diff --git a/mocha/nodes.py b/mocha/nodes.py new file mode 100644 index 0000000..2f46614 --- /dev/null +++ b/mocha/nodes.py @@ -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", + } \ No newline at end of file diff --git a/nodes_sampler.py b/nodes_sampler.py index 63596fc..4bd5f7f 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -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 @@ -415,6 +415,12 @@ class WanVideoSampler: log.info(f"RecamMaster camera embed shape: {camera_embed.shape}") 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) @@ -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 @@ -1188,6 +1206,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 \ @@ -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() diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 525d57d..c2e32b6 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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)