From 75190b756ab0302bc4b6a895415a74722b2327cd Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 30 Mar 2025 19:07:04 +0300 Subject: [PATCH] Add Mobius looping option https://github.com/YisuiTT/Mobius/ --- .../pipelines/pipeline_hunyuan_video.py | 26 +++++++++++++++++-- nodes.py | 24 ++++++++++++++++- 2 files changed, 47 insertions(+), 3 deletions(-) diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 1ee7db4..f612ad1 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -454,6 +454,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): image_cond_latents: Optional[torch.Tensor] = None, riflex_freq_index: Optional[int] = None, i2v_stability=True, + loop_args: Optional[Dict] = None, **kwargs, ): r""" @@ -689,6 +690,15 @@ class HunyuanVideoPipeline(DiffusionPipeline): callback = prepare_callback(self.comfy_model, num_inference_steps) #print(self.scheduler.sigmas) + + latent_shift_loop = False + if loop_args is not None: + latent_shift_loop = True + is_looped = True + latent_skip = loop_args["shift_skip"] + latent_shift_start_percent = loop_args["start_percent"] + latent_shift_end_percent = loop_args["end_percent"] + shift_idx = 0 logger.info(f"Sampling {video_length} frames in {latents.shape[2]} latents at {width}x{height} with {len(timesteps)} inference steps") @@ -698,9 +708,11 @@ class HunyuanVideoPipeline(DiffusionPipeline): if self.interrupt: continue + current_step_percentage = i / len(timesteps) + if image_cond_latents is not None and i2v_condition_type == "token_replace": latents = torch.concat([original_image_latents, latents[:, :, 1:, :, :]], dim=2) - + latent_model_input = latents input_prompt_embeds = prompt_embeds #input_prompt_mask = prompt_mask @@ -708,7 +720,12 @@ class HunyuanVideoPipeline(DiffusionPipeline): cfg_enabled = False stg_enabled = False - current_step_percentage = i / len(timesteps) + ### latent shift + if latent_shift_loop: + if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent: + latent_model_input = torch.cat([latent_model_input[:, :, shift_idx:]] + [latent_model_input[:, :, :shift_idx]], dim=2) + + if self.do_spatio_temporal_guidance: if stg_start_percent <= current_step_percentage <= stg_end_percent: stg_enabled = True @@ -904,6 +921,11 @@ class HunyuanVideoPipeline(DiffusionPipeline): noise_pred = noise_pred_text + self._stg_scale * ( noise_pred_text - noise_pred_perturb ) + if latent_shift_loop: + #reverse latent shift + if latent_shift_start_percent <= current_step_percentage <= latent_shift_end_percent: + noise_pred = torch.cat([noise_pred[:, :, latent_video_length - shift_idx:]] + [noise_pred[:, :, :latent_video_length - shift_idx]], dim=2) + shift_idx = (shift_idx + latent_skip) % latent_video_length # compute the previous noisy sample x_t -> x_t-1 if image_cond_latents is not None and i2v_condition_type == "token_replace": diff --git a/nodes.py b/nodes.py index 8833aae..3f6b601 100644 --- a/nodes.py +++ b/nodes.py @@ -1266,6 +1266,26 @@ class HyVideoContextOptions: } return (context_options,) + +class HyVideoLoopArgs: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "shift_skip": ("INT", {"default": 6, "min": 0, "tooltip": "Skip step of latent shift"}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the looping effect"}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the looping effect"}), + }, + } + + RETURN_TYPES = ("LOOPARGS", ) + RETURN_NAMES = ("loop_args",) + FUNCTION = "process" + CATEGORY = "HunyuanVideoWrapper" + DESCRIPTION = "Looping through latent shift as shown in https://github.com/YisuiTT/Mobius/" + + def process(self, **kwargs): + return (kwargs,) + #region Sampler class HyVideoSampler: @classmethod @@ -1298,6 +1318,7 @@ class HyVideoSampler: }), "riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 4. Allows for new frames to be generated after 129 without looping"}), "i2v_mode": (["stability", "dynamic"], {"default": "dynamic", "tooltip": "I2V mode for image2video process"}), + "loop_args": ("LOOPARGS", ), } } @@ -1308,7 +1329,7 @@ class HyVideoSampler: def process(self, model, hyvid_embeds, flow_shift, steps, embedded_guidance_scale, seed, width, height, num_frames, samples=None, denoise_strength=1.0, force_offload=True, stg_args=None, context_options=None, feta_args=None, - teacache_args=None, scheduler=None, image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability"): + teacache_args=None, scheduler=None, image_cond_latents=None, riflex_freq_index=0, i2v_mode="stability", loop_args=None): model = model.model device = mm.get_torch_device() @@ -1462,6 +1483,7 @@ class HyVideoSampler: image_cond_latents = image_cond_latents["samples"] * VAE_SCALING_FACTOR if image_cond_latents is not None else None, riflex_freq_index = riflex_freq_index, i2v_stability = i2v_stability, + loop_args = loop_args, ) print_memory(device)