@@ -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":
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user