diff --git a/control/control.py b/control/control.py index df53914..b473b3e 100644 --- a/control/control.py +++ b/control/control.py @@ -306,6 +306,13 @@ class SparseCtrlAdvanced(ControlNetAdvanced): control = self.control_model(x=x_noisy.to(dtype), hint=self.cond_hint, timesteps=timestep.float(), context=context.to(dtype), y=y) return self.control_merge(None, control, control_prev, output_dtype) + def pre_run_advanced(self, model, percent_to_timestep_function): + super().pre_run_advanced(model, percent_to_timestep_function) + if self.control_model.motion_holder is not None: + self.control_model.motion_holder.motion_wrapper.reset() + self.control_model.motion_holder.motion_wrapper.set_strength(self.sparse_settings.motion_strength) + self.control_model.motion_holder.motion_wrapper.set_scale_multiplier(self.sparse_settings.motion_scale) + def copy(self): c = SparseCtrlAdvanced(self.control_model, self.timestep_keyframes, self.sparse_settings, self.global_average_pooling, self.device, self.load_device, self.manual_cast_dtype) self.copy_to(c) diff --git a/control/control_sparsectrl.py b/control/control_sparsectrl.py index 717e745..1abb8cd 100644 --- a/control/control_sparsectrl.py +++ b/control/control_sparsectrl.py @@ -80,9 +80,11 @@ class SparseControlNet(ControlNetCLDM): class SparseSettings: - def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True): + def __init__(self, sparse_method: 'SparseMethod', use_motion: bool=True, motion_strength=1.0, motion_scale=1.0): self.sparse_method = sparse_method self.use_motion = use_motion + self.motion_strength = motion_strength + self.motion_scale = motion_scale @classmethod def default(cls): @@ -318,18 +320,32 @@ class SparseCtrlMotionWrapper(nn.Module): self.mid_block.set_video_length(video_length, full_length) def set_scale_multiplier(self, multiplier: Union[float, None]): - for block in self.down_blocks: - block.set_scale_multiplier(multiplier) - for block in self.up_blocks: - block.set_scale_multiplier(multiplier) + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_scale_multiplier(multiplier) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_scale_multiplier(multiplier) if self.mid_block is not None: self.mid_block.set_scale_multiplier(multiplier) + def set_strength(self, strength: float): + if self.down_blocks is not None: + for block in self.down_blocks: + block.set_strength(strength) + if self.up_blocks is not None: + for block in self.up_blocks: + block.set_strength(strength) + if self.mid_block is not None: + self.mid_block.set_strength(strength) + def reset_temp_vars(self): - for block in self.down_blocks: - block.reset_temp_vars() - for block in self.up_blocks: - block.reset_temp_vars() + if self.down_blocks is not None: + for block in self.down_blocks: + block.reset_temp_vars() + if self.up_blocks is not None: + for block in self.up_blocks: + block.reset_temp_vars() if self.mid_block is not None: self.mid_block.reset_temp_vars() @@ -375,6 +391,10 @@ class MotionModule(nn.Module): for motion_module in self.motion_modules: motion_module.set_sub_idxs(sub_idxs) + def set_strength(self, strength: float): + for motion_module in self.motion_modules: + motion_module.set_strength(strength) + def reset_temp_vars(self): for motion_module in self.motion_modules: motion_module.reset_temp_vars() @@ -399,7 +419,7 @@ class VanillaTemporalModule(nn.Module): zero_initialize=True, ): super().__init__() - + self.strength = 1.0 self.temporal_transformer = TemporalTransformer3DModel( in_channels=in_channels, num_attention_heads=num_attention_heads, @@ -430,11 +450,22 @@ class VanillaTemporalModule(nn.Module): def set_sub_idxs(self, sub_idxs: list[int]): self.temporal_transformer.set_sub_idxs(sub_idxs) + def set_strength(self, strength: float): + self.strength = strength + def reset_temp_vars(self): + self.set_strength(1.0) self.temporal_transformer.reset_temp_vars() def forward(self, input_tensor, encoder_hidden_states=None, attention_mask=None): - return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + if math.isclose(self.strength, 1.0): + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + elif math.isclose(self.strength, 0.0): + return input_tensor + elif self.strength > 1.0: + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength + else: + return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*self.strength + input_tensor*(1.0-self.strength) class TemporalTransformer3DModel(nn.Module): diff --git a/control/nodes_sparsectrl.py b/control/nodes_sparsectrl.py index 7759013..2faae66 100644 --- a/control/nodes_sparsectrl.py +++ b/control/nodes_sparsectrl.py @@ -14,6 +14,8 @@ class SparseCtrlLoaderAdvanced: "required": { "control_net_name": (folder_paths.get_filename_list("controlnet"), ), "use_motion": ("BOOLEAN", {"default": True}, ), + "motion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "motion_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), }, "optional": { "sparse_method": ("SPARSE_METHOD", ), @@ -26,9 +28,9 @@ class SparseCtrlLoaderAdvanced: CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl" - def load_controlnet(self, control_net_name: str, use_motion: bool, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): + def load_controlnet(self, control_net_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None): controlnet_path = folder_paths.get_full_path("controlnet", control_net_name) - sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion) + sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale) controlnet = load_sparsectrl(controlnet_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings) return (controlnet,)