Fixed ControlLLLite issues due to deepcopy behavior within ModelPatcher patches

This commit is contained in:
Jedrzej Kosinski
2024-02-02 02:13:35 -06:00
parent 04bbe21589
commit be2da1c086
2 changed files with 6 additions and 8 deletions
+6 -3
View File
@@ -301,7 +301,7 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
def pre_run_advanced(self, *args, **kwargs):
AdvancedControlBase.pre_run_advanced(self, *args, **kwargs)
self.patch.control = self
self.patch.set_control(self)
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int):
# normal ControlNet stuff
@@ -341,6 +341,10 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
self.copy_to(c)
self.copy_to_advanced(c)
return c
# deepcopy needs to properly keep track of objects to work between model.clone calls!
def __deepcopy__(self, *args, **kwargs):
return self
# def get_models(self):
# # get_models is called once at the start of every KSampler run - use to reset already_patched status
@@ -358,7 +362,6 @@ def load_controlnet(ckpt_path, timestep_keyframe: TimestepKeyframeGroup=None, mo
for key in controlnet_data:
# LLLLite check
if "lllite" in key:
logger.info("ControlLLLite controlnet!")
controlnet_type = ControlWeightType.CONTROLLLLITE
break
# SparseCtrl check
@@ -596,7 +599,7 @@ def load_controllllite(ckpt_path: str, controlnet_data: dict[str, Tensor]=None,
if len(modules) == 1:
module.is_first = True
logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules")
#logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules")
patch = LLLitePatch(modules=modules)
control = ControlLLLiteAdvanced(patch=patch, timestep_keyframes=timestep_keyframe)
-5
View File
@@ -62,11 +62,6 @@ class LLLitePatch:
module_pfx_to_k = module_pfx + "_to_k"
module_pfx_to_v = module_pfx + "_to_v"
# if masks present, get masks with same dims as attention
# if q.shape != k.shape or q.shape != v.shape:
# logger.warn(f"mismatch!!! q:{q.shape}, k:{k.shape}, v:{v.shape}")
#logger.warn(f"{q.shape}")
if module_pfx_to_q in self.modules:
q = q + self.modules[module_pfx_to_q](q, self.control)
if module_pfx_to_k in self.modules: