Add motion_strength and motion_scale to Load SparseCtrl node
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user