From 6cd7f2296f4240ce10ccb31f53ca7325d0848125 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 17 Mar 2025 19:45:52 +0200 Subject: [PATCH] Add Moebius latent shift looping https://github.com/YisuiTT/Mobius/ --- nodes.py | 41 +++++++++++++++++++++++++++++++++++++++-- 1 file changed, 39 insertions(+), 2 deletions(-) diff --git a/nodes.py b/nodes.py index 6c1c7b0..9f56447 100644 --- a/nodes.py +++ b/nodes.py @@ -1216,6 +1216,23 @@ class WanVideoFlowEdit: def process(self, **kwargs): return (kwargs,) +class WanVideoLoopArgs: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "shift_skip": ("INT", {"default": 6, "min": 0}), + }, + } + + RETURN_TYPES = ("LOOPARGS", ) + RETURN_NAMES = ("loop_args",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Looping through latent shift as shown in https://github.com/YisuiTT/Mobius/" + + def process(self, **kwargs): + return (kwargs,) + class WanVideoSampler: @classmethod def INPUT_TYPES(s): @@ -1247,6 +1264,7 @@ class WanVideoSampler: "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"}), + "loop_args": ("LOOPARGS", ), } } @@ -1257,7 +1275,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, rope_function="default"): + teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None): #assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache." patcher = model model = model.model @@ -1650,7 +1668,13 @@ class WanVideoSampler: thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device) masks = mask.repeat(len(timesteps), 1, 1, 1, 1).to(device) masks = masks > thresholds - + + if loop_args is not None: + latent_shift_loop = True + is_looped = True + latent_skip = loop_args["shift_skip"] + shift_idx = 0 + #main loop start for idx, t in enumerate(tqdm(timesteps)): if flowedit_args is not None: if idx < skip_steps: @@ -1669,6 +1693,11 @@ class WanVideoSampler: # end diff diff latent_model_input = latent.to(device) + + ### latent shift + if latent_shift_loop: + latent_model_input = torch.cat([latent_model_input[:, shift_idx:]] + [latent_model_input[:, :shift_idx]], dim=1) + timestep = torch.tensor([t]).to(device) current_step_percentage = idx / len(timesteps) @@ -1878,6 +1907,12 @@ class WanVideoSampler: text_embeds["negative_prompt_embeds"], timestep, idx, image_cond, clip_fea, teacache_state=self.teacache_state) + + if latent_shift_loop: + #reverse latent shift + noise_pred = torch.cat([noise_pred[:, latent_video_length - shift_idx:]] + [noise_pred[:, :latent_video_length - shift_idx]], dim=1) + shift_idx = (shift_idx + latent_skip) % latent_video_length + if flowedit_args is None: latent = latent.to(intermediate_device) @@ -2196,6 +2231,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoControlEmbeds": WanVideoControlEmbeds, "WanVideoSLG": WanVideoSLG, "WanVideoTinyVAELoader": WanVideoTinyVAELoader, + "WanVideoLoopArgs": WanVideoLoopArgs } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -2223,4 +2259,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoControlEmbeds": "WanVideo Control Embeds", "WanVideoSLG": "WanVideo SLG", "WanVideoTinyVAELoader": "WanVideo Tiny VAE Loader", + "WanVideoLoopArgs": "WanVideo Loop Args" }