Pad the input by one latent when using I2V

Seems better as we drop the first latent at end anyway
This commit is contained in:
kijai
2025-03-06 20:09:51 +02:00
parent d8bd580842
commit 16b2c11d2b
@@ -236,11 +236,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
context_size=None,
context_overlap=None,
leapfusion_img2vid=False,
i2v_mask=None,
official_i2v=False,
image_cond_latents=None,
):
#if i2v_mask is not None:
# num_channels_latents = (num_channels_latents - 1) // 2
shape = (
batch_size,
num_channels_latents,
@@ -288,11 +286,16 @@ class HunyuanVideoPipeline(DiffusionPipeline):
# apply shuffled indexes
#print("place_idx:", place_idx, "delta:", delta, "list_idx:", list_idx)
noise[:, :, place_idx:place_idx + delta, :, :] = noise[:, :, list_idx, :, :]
if i2v_mask is not None:
print("i2v_mask shape:", i2v_mask.shape)
if image_cond_latents.shape[2] == 1:
image_cond_latents = image_cond_latents.repeat(1, 1, video_length, 1, 1)
i2v_mask = None
if image_cond_latents is not None:
if image_cond_latents.shape[2] == 1:
padding = torch.zeros(shape, device=device)
padding[:, :, 0:1, :, :] = image_cond_latents
image_cond_latents = padding
if official_i2v:
# Create mask
i2v_mask = torch.zeros(shape[0], 1, shape[2], shape[3], shape[4], device=device)
i2v_mask[:, :, 0, ...] = 1.0
t = torch.tensor([0.999]).to(device=device)
latents = noise * t + image_cond_latents * (1 - t)
latents = latents.to(dtype=self.base_dtype)
@@ -321,7 +324,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if hasattr(self.scheduler, "init_noise_sigma"):
# scale the initial noise by the standard deviation required by the scheduler
latents = latents * self.scheduler.init_noise_sigma
return latents, timesteps
return latents, timesteps, i2v_mask, image_cond_latents
# Copied from diffusers.pipelines.latent_consistency_models.pipeline_latent_consistency_text2img.LatentConsistencyModelPipeline.get_guidance_scale_embedding
def get_guidance_scale_embedding(
@@ -578,6 +581,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latent_video_length = (video_length - 1) // 4 + 1
official_i2v = False
if self.transformer.in_channels == 33:
official_i2v = True
latent_video_length += 1
if feta_args is not None:
set_enhance_weight(feta_args["weight"])
feta_start_percent = feta_args["start_percent"]
@@ -586,29 +594,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else:
disable_enhance()
i2v_mask = None
image_latents = None
if image_cond_latents is not None:
# Expand to video length and zero-pad remaining frames
image_latents = torch.zeros(
(batch_size, 16, latent_video_length, height//8, width//8),
device=device,
dtype=self.base_dtype
)
image_latents[:, :, 0:1, ...] = image_cond_latents
# Create mask
i2v_mask = torch.zeros(
batch_size, 1, latent_video_length, height//8, width//8,
device=device
)
i2v_mask[:, :, 0, ...] = 1.0
print("i2v_mask shape:", i2v_mask.shape)
print("image_cond_latents shape:", image_cond_latents.shape)
print("image_latents shape:", image_latents.shape)
# context windows
use_context_schedule = False
freenoise = False
@@ -647,7 +632,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
# 5. Prepare latent variables
#num_channels_latents = self.transformer.config.in_channels
num_channels_latents = 16
latents, timesteps = self.prepare_latents(
latents, timesteps, i2v_mask, image_cond_latents = self.prepare_latents(
batch_size * num_videos_per_prompt,
num_channels_latents,
num_inference_steps,
@@ -663,8 +648,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
context_size=context_frames,
context_overlap=context_overlap,
leapfusion_img2vid=leapfusion_img2vid,
i2v_mask=i2v_mask,
image_cond_latents=image_latents,
official_i2v=official_i2v,
image_cond_latents=image_cond_latents,
)
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
@@ -742,7 +727,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if image_cond_latents is not None:
latent_image_input = (
torch.cat([image_latents] * 2) if cfg_enabled else image_latents
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
@@ -911,6 +896,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else:
comfy_pbar.update(1)
if leapfusion_img2vid or i2v_mask is not None:
if leapfusion_img2vid or official_i2v:
latents = latents[:, :, 1:, :, :]
return latents