@@ -787,6 +787,7 @@ def __call__(
787787 num_inference_steps : int = 20 ,
788788 mode : str = "animate" ,
789789 prev_segment_conditioning_frames : int = 1 ,
790+ motion_encode_batch_size : Optional [int ] = None ,
790791 guidance_scale : float = 1.0 ,
791792 num_videos_per_prompt : Optional [int ] = 1 ,
792793 generator : Optional [Union [torch .Generator , List [torch .Generator ]]] = None ,
@@ -832,6 +833,10 @@ def __call__(
832833 prev_segment_conditioning_frames (`int`, defaults to `1`):
833834 The number of frames from the previous video segment to be used for temporal guidance. Recommended to
834835 be 1 or 5. In general, should be 4N + 1, where N is a non-negative integer.
836+ motion_encode_batch_size (`int`, *optional*):
837+ The batch size for batched encoding of the face video via the motion encoder. This allows trading off
838+ inference speed for lower memory usage by setting a smaller batch size. Will default to
839+ `self.transformer.config.motion_encoder_batch_size` if not set.
835840 height (`int`, defaults to `720`):
836841 The height of the generated video.
837842 width (`int`, defaults to `1280`):
@@ -1127,6 +1132,7 @@ def __call__(
11271132 encoder_hidden_states_image = image_embeds ,
11281133 pose_hidden_states = pose_latents ,
11291134 face_pixel_values = face_video_segment ,
1135+ motion_encode_batch_size = motion_encode_batch_size ,
11301136 attention_kwargs = attention_kwargs ,
11311137 return_dict = False ,
11321138 )[0 ]
@@ -1142,6 +1148,7 @@ def __call__(
11421148 encoder_hidden_states_image = image_embeds ,
11431149 pose_hidden_states = pose_latents ,
11441150 face_pixel_values = face_pixel_values_uncond ,
1151+ motion_encode_batch_size = motion_encode_batch_size ,
11451152 attention_kwargs = attention_kwargs ,
11461153 return_dict = False ,
11471154 )[0 ]
0 commit comments