From d4d93377e57bc3c4358931bd8dcd60e73de91b29 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 7 Mar 2025 22:38:04 +0200 Subject: [PATCH] fix stability mode for the old model, add vae dist mode selection --- hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py | 8 +++++--- hyvideo/modules/models.py | 1 - nodes.py | 9 +++++++-- 3 files changed, 12 insertions(+), 6 deletions(-) diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 7f30bd0..fe93443 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -320,6 +320,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): elif image_cond_latents is not None and i2v_stability: if image_cond_latents.shape[2] == 1: img_latents = image_cond_latents.repeat(1, 1, video_length, 1, 1) + else: + img_latents = image_cond_latents t = torch.tensor([0.999]).to(device=device) latents = noise * t + img_latents * (1 - t) latents = latents.to(dtype=self.base_dtype) @@ -728,8 +730,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): t_expand = t.repeat(latent_model_input.shape[0]) - if leapfusion_img2vid: - latent_model_input[:, :, [0,], :, :] = original_latents[:, :, [0,], :, :].to(latent_model_input) + #if leapfusion_img2vid: + # latent_model_input[:, :, [0,], :, :] = original_latents[:, :, [0,], :, :].to(latent_model_input) if image_cond_latents is not None and not use_context_schedule: if i2v_condition_type == "latent_concat": @@ -737,7 +739,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): i2v_mask = torch.cat([i2v_mask] * 2) if cfg_enabled else i2v_mask latent_image_input = torch.cat([latent_image_input, i2v_mask], dim=1) latent_model_input = torch.cat([latent_model_input, latent_image_input], dim=1) - elif i2v_condition_type == "token_replace": + elif i2v_condition_type == "token_replace" or leapfusion_img2vid: latent_image_input = (torch.cat([original_image_latents] * 2) if cfg_enabled else original_image_latents) latent_model_input = torch.cat([latent_image_input, latent_model_input[:, :, 1:, :, :]], dim=2) else: diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 1f16ca3..cb61c84 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -201,7 +201,6 @@ class MMDoubleStreamBlock(nn.Module): first_frame_token_num: int = None, condition_type: str = None, ) -> Tuple[torch.Tensor, torch.Tensor]: - if condition_type == "token_replace": img_mod1, token_replace_img_mod1 = self.img_mod(vec, condition_type=condition_type, \ token_replace_vec=token_replace_vec) diff --git a/nodes.py b/nodes.py index e588c8d..ce95013 100644 --- a/nodes.py +++ b/nodes.py @@ -1489,6 +1489,7 @@ class HyVideoEncode: "optional": { "noise_aug_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Strength of noise augmentation, helpful for leapfusion I2V where some noise can add motion and give sharper results"}), "latent_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional latent multiplier, helpful for leapfusion I2V where lower values allow for more motion"}), + "latent_dist": (["sample", "mode"], {"default": "sample", "tooltip": "Sampling mode for the VAE, sample uses the latent distribution, mode uses the mode of the latent distribution"}), } } @@ -1497,7 +1498,8 @@ class HyVideoEncode: FUNCTION = "encode" CATEGORY = "HunyuanVideoWrapper" - def encode(self, vae, image, enable_vae_tiling, temporal_tiling_sample_size, auto_tile_size, spatial_tile_sample_min_size, noise_aug_strength=0.0, latent_strength=1.0): + def encode(self, vae, image, enable_vae_tiling, temporal_tiling_sample_size, auto_tile_size, + spatial_tile_sample_min_size, noise_aug_strength=0.0, latent_strength=1.0, latent_dist="sample"): device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -1523,7 +1525,10 @@ class HyVideoEncode: if enable_vae_tiling: vae.enable_tiling() - latents = vae.encode(image).latent_dist.sample(generator) + if latent_dist == "sample": + latents = vae.encode(image).latent_dist.sample(generator) + elif latent_dist == "mode": + latents = vae.encode(image).latent_dist.mode() if latent_strength != 1.0: latents *= latent_strength #latents = latents * vae.config.scaling_factor