Add Mobius looping option

https://github.com/YisuiTT/Mobius/
This commit is contained in:
kijai
2025-03-30 19:07:04 +03:00
parent 1b4164f5b3
commit 75190b756a
2 changed files with 47 additions and 3 deletions
@@ -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":
+23 -1
View File
@@ -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)