From b78cc7f5c7fed1f966fdd0c9e577a61edeae25d5 Mon Sep 17 00:00:00 2001 From: Zoybzo <61582821+Zoybzo@users.noreply.github.com> Date: Tue, 13 May 2025 16:23:39 +0800 Subject: [PATCH] Fixed the single image infer error --- models/pipeline_a2.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/models/pipeline_a2.py b/models/pipeline_a2.py index ee89115..7ba08e2 100644 --- a/models/pipeline_a2.py +++ b/models/pipeline_a2.py @@ -428,7 +428,7 @@ def prepare_latents( mask_lat_size[:, :, list(range(1, num_frames))] = 0 first_frame_mask = mask_lat_size[:, :, 0:1] if len(image_vae) == 1: - first_frame_mask = torch.repeat_interleave(first_frame_mask, dim=2, repeats=self.vae_scale_factor_temporal+4) + first_frame_mask = torch.repeat_interleave(first_frame_mask, dim=2, repeats=self.vae_scale_factor_temporal) mask_lat_size = torch.concat([first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2) elif len(image_vae) == 2: first_frame_mask = torch.repeat_interleave(first_frame_mask, dim=2, repeats=self.vae_scale_factor_temporal+4) @@ -737,4 +737,4 @@ def __call__( if not return_dict: return (video,) - return WanPipelineOutput(frames=video) \ No newline at end of file + return WanPipelineOutput(frames=video)