Add motion_strength and motion_scale to Load SparseCtrl node

This commit is contained in:
Jedrzej Kosinski
2023-12-21 23:35:36 -06:00
parent 2fc43f458b
commit 2ca56093ff
3 changed files with 53 additions and 13 deletions
+7
View File
@@ -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)
+42 -11
View File
@@ -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):
+4 -2
View File
@@ -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,)