From b66fe529b1ad889fa860bd3dff256e0335f4939e Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 3 Jul 2024 17:58:03 +0300 Subject: [PATCH] Add strength controls --- mimicmotion/modules/unet.py | 4 ++- mimicmotion/pipelines/pipeline_mimicmotion.py | 36 +++++++++++++++++-- nodes.py | 12 +++++-- 3 files changed, 46 insertions(+), 6 deletions(-) diff --git a/mimicmotion/modules/unet.py b/mimicmotion/modules/unet.py index 4e91bd0..b46bd98 100644 --- a/mimicmotion/modules/unet.py +++ b/mimicmotion/modules/unet.py @@ -368,6 +368,7 @@ class UNetSpatioTemporalConditionModel(ModelMixin, ConfigMixin, UNet2DConditionL encoder_hidden_states: torch.Tensor, added_time_ids: torch.Tensor, pose_latents: torch.Tensor = None, + pose_strength: float = 1.0, image_only_indicator: bool = False, return_dict: bool = True, ) -> Union[UNetSpatioTemporalConditionOutput, Tuple]: @@ -437,11 +438,12 @@ class UNetSpatioTemporalConditionModel(ModelMixin, ConfigMixin, UNet2DConditionL emb = emb.repeat_interleave(num_frames, dim=0) # encoder_hidden_states: [batch, 1, channels] -> [batch * frames, 1, channels] encoder_hidden_states = encoder_hidden_states.repeat_interleave(num_frames, dim=0) + encoder_hidden_states = encoder_hidden_states # 2. pre-process sample = self.conv_in(sample) if pose_latents is not None: - sample = sample + pose_latents + sample = sample + pose_latents * pose_strength image_only_indicator = torch.ones(batch_size, num_frames, dtype=sample.dtype, device=sample.device) \ if image_only_indicator else torch.zeros(batch_size, num_frames, dtype=sample.dtype, device=sample.device) diff --git a/mimicmotion/pipelines/pipeline_mimicmotion.py b/mimicmotion/pipelines/pipeline_mimicmotion.py index 3f5dbc2..40245ae 100644 --- a/mimicmotion/pipelines/pipeline_mimicmotion.py +++ b/mimicmotion/pipelines/pipeline_mimicmotion.py @@ -124,7 +124,8 @@ class MimicMotionPipeline(DiffusionPipeline): image: PipelineImageInput, device: Union[str, torch.device], num_videos_per_prompt: int, - do_classifier_free_guidance: bool): + do_classifier_free_guidance: bool, + image_embed_strength: float = 1.0): dtype = next(self.image_encoder.parameters()).dtype # if not isinstance(image, torch.Tensor): @@ -160,6 +161,7 @@ class MimicMotionPipeline(DiffusionPipeline): bs_embed, seq_len, _ = image_embeddings.shape image_embeddings = image_embeddings.repeat(1, num_videos_per_prompt, 1) image_embeddings = image_embeddings.view(bs_embed * num_videos_per_prompt, seq_len, -1) + image_embeddings = image_embeddings * image_embed_strength if do_classifier_free_guidance: negative_image_embeddings = torch.zeros_like(image_embeddings) @@ -169,6 +171,8 @@ class MimicMotionPipeline(DiffusionPipeline): # Here we concatenate the unconditional and text embeddings into a single batch # to avoid doing two forward passes image_embeddings = torch.cat([negative_image_embeddings, image_embeddings]) + + return image_embeddings @@ -179,8 +183,10 @@ class MimicMotionPipeline(DiffusionPipeline): ): # Get latents_pose pose_latents = self.pose_net(pose_image) + print(pose_latents.shape) if do_classifier_free_guidance: + print("doing classifier free guidance") negative_pose_latents = torch.zeros_like(pose_latents) # For classifier free guidance, we need to do two forward passes. @@ -371,6 +377,10 @@ class MimicMotionPipeline(DiffusionPipeline): self, image: Union[PIL.Image.Image, List[PIL.Image.Image], torch.FloatTensor], image_pose: Union[torch.FloatTensor], + pose_strength: float = 1.0, + pose_start_percent: float = 0.0, + pose_end_percent: float = 1.0, + image_embed_strength: float = 1.0, height: int = 576, width: int = 1024, num_frames: Optional[int] = None, @@ -508,7 +518,7 @@ class MimicMotionPipeline(DiffusionPipeline): self._guidance_scale = max_guidance_scale # 3. Encode input image - image_embeddings = self._encode_image(image, device, num_videos_per_prompt, self.do_classifier_free_guidance) + image_embeddings = self._encode_image(image, device, num_videos_per_prompt, self.do_classifier_free_guidance, image_embed_strength=image_embed_strength) # NOTE: Stable Diffusion Video was conditioned on fps - 1, which # is why it is reduced here. @@ -589,6 +599,13 @@ class MimicMotionPipeline(DiffusionPipeline): self.unet.to(device) self._num_timesteps = len(timesteps) + + # Calculate the actual start and end steps based on percentages + start_step_index = round(self._num_timesteps * pose_start_percent) + end_step_index = round(self._num_timesteps * pose_end_percent) + + print(f"start_step_index: {start_step_index}, end_step_index: {end_step_index}") + pose_latents = einops.rearrange(pose_latents, '(b f) c h w -> b f c h w', f=num_frames) indices = [[0, *range(i + 1, min(i + tile_size, num_frames))] for i in range(0, num_frames - tile_size + 1, tile_size - tile_overlap)] @@ -610,13 +627,26 @@ class MimicMotionPipeline(DiffusionPipeline): # image_pose = pixel_values_pose[:, frame_start:frame_start + self.num_frames, ...] weight = (torch.arange(tile_size, device=device) + 0.5) * 2. / tile_size weight = torch.minimum(weight, 2 - weight) + for idx in indices: + # Check if the current timestep is within the start and end step range + if start_step_index <= i <= end_step_index: + # Apply pose_latents as currently done + print(f"Applying pose on step {i}") + pose_latents_to_use = pose_latents[:, idx].flatten(0, 1) + else: + print(f"Not applying pose on step {i}") + # Apply an alternative if pose_latents should not be used outside this range + # This could be zeros, or any other placeholder logic you define. + pose_latents_to_use = torch.zeros_like(pose_latents[:, idx].flatten(0, 1)) + _noise_pred = self.unet( latent_model_input[:, idx], t, encoder_hidden_states=image_embeddings, added_time_ids=added_time_ids, - pose_latents=pose_latents[:, idx].flatten(0, 1), + pose_latents=pose_latents_to_use, + pose_strength=pose_strength, image_only_indicator=image_only_indicator, return_dict=False, )[0] diff --git a/nodes.py b/nodes.py index 3dc25d7..216d10a 100644 --- a/nodes.py +++ b/nodes.py @@ -242,6 +242,10 @@ class MimicMotionSampler: }, "optional": { "optional_scheduler": ("DIFFUSERS_SCHEDULER",), + "pose_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "pose_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "pose_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "image_embed_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), } } @@ -251,7 +255,7 @@ class MimicMotionSampler: CATEGORY = "MimicMotionWrapper" def process(self, mimic_pipeline, ref_image, pose_images, cfg_min, cfg_max, steps, seed, noise_aug_strength, fps, keep_model_loaded, - context_size, context_overlap, optional_scheduler=None): + context_size, context_overlap, optional_scheduler=None, pose_strength=1.0, image_embed_strength=1.0, pose_start_percent=0.0, pose_end_percent=1.0): device = mm.get_torch_device() offload_device = mm.unet_offload_device() mm.unload_all_models() @@ -307,7 +311,11 @@ class MimicMotionSampler: decode_chunk_size=4, output_type="latent", device=device, - sigmas=sigmas + sigmas=sigmas, + pose_strength=pose_strength, + pose_start_percent=pose_start_percent, + pose_end_percent=pose_end_percent, + image_embed_strength=image_embed_strength ).frames if not keep_model_loaded: