From e4b52667d6c7735824974822924626e67400282c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 6 Mar 2025 23:09:37 +0200 Subject: [PATCH] Update pipeline_hunyuan_video.py --- .../pipelines/pipeline_hunyuan_video.py | 32 +++++++++++++++---- 1 file changed, 26 insertions(+), 6 deletions(-) diff --git a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py index 03acdfe..5c40f56 100644 --- a/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py +++ b/hyvideo/diffusion/pipelines/pipeline_hunyuan_video.py @@ -235,7 +235,6 @@ class HunyuanVideoPipeline(DiffusionPipeline): freenoise=False, context_size=None, context_overlap=None, - leapfusion_img2vid=False, official_i2v=False, image_cond_latents=None, ): @@ -647,7 +646,6 @@ class HunyuanVideoPipeline(DiffusionPipeline): freenoise=freenoise, context_size=context_frames, context_overlap=context_overlap, - leapfusion_img2vid=leapfusion_img2vid, official_i2v=official_i2v, image_cond_latents=image_cond_latents, ) @@ -725,7 +723,7 @@ class HunyuanVideoPipeline(DiffusionPipeline): if leapfusion_img2vid: latent_model_input[:, :, [0,], :, :] = original_latents[:, :, [0,], :, :].to(latent_model_input) - if image_cond_latents is not None: + if image_cond_latents is not None and not use_context_schedule: latent_image_input = ( torch.cat([image_cond_latents] * 2) if cfg_enabled else image_cond_latents ) @@ -749,13 +747,33 @@ class HunyuanVideoPipeline(DiffusionPipeline): guidance_expand = None if use_context_schedule: - counter = torch.zeros_like(latent_model_input) - noise_pred = torch.zeros_like(latent_model_input) + counter = torch.zeros_like(latent_model_input)[:, :16] + noise_pred = torch.zeros_like(latent_model_input)[:, :16] + print("noise_pred", noise_pred.shape) + print("counter", counter.shape) context_queue = list(context( i, num_inference_steps, latents.shape[2], context_frames, context_stride, context_overlap, )) + + if image_cond_latents is not None: + latent_image_input = ( + torch.cat([image_cond_latents] * 2) if cfg_enabled else image_cond_latents + ) + if i2v_mask is not None: + i2v_mask = torch.cat([i2v_mask] * 2) if cfg_enabled else i2v_mask + for c in context_queue: - partial_latent_model_input = latent_model_input[:, :, c, :, :] + partial_latent_model_input = latent_model_input[:, :, c] + if i2v_mask is not None: + #doesn't work properly + new_mask = torch.zeros_like(i2v_mask) + new_mask[..., 0, :, :] = 1.0 + new_image_input = torch.zeros_like(latent_image_input) + new_image_input[..., 0, :, :] = latent_image_input[..., 0, :, :] + + partial_latent_image_input = torch.cat([new_image_input[..., c, :, :], new_mask[..., c, :, :]], dim=1) + partial_latent_model_input = torch.cat([partial_latent_model_input, partial_latent_image_input], dim=1) + #print("partial_latent_model_input", partial_latent_model_input.shape) with torch.autocast( device_type="cuda", dtype=self.base_dtype, enabled=True): @@ -773,6 +791,8 @@ class HunyuanVideoPipeline(DiffusionPipeline): return_dict=True, )["x"] window_mask = torch.ones_like(noise_pred_context) + + # Apply left-side blending for all except first chunk if min(c) > 0: ramp_up = torch.linspace(0, 1, context_overlap, device=noise_pred_context.device)