Make SparseCtrl work with new ComfyUI ModelPatcher updates

This commit is contained in:
Jedrzej Kosinski
2024-08-20 05:09:30 -05:00
parent 85d4970cae
commit dd7e29e5c2
+10 -2
View File
@@ -114,9 +114,13 @@ class SparseModelPatcher(ModelPatcher):
self.model: SparseControlNet
super().__init__(*args, **kwargs)
def patch_model_lowvram(self, device_to=None, *args, **kwargs):
patched_model = super().patch_model_lowvram(device_to, *args, **kwargs)
def load(self, device_to=None, lowvram_model_memory=0, *args, **kwargs):
to_return = super().load(device_to=device_to, lowvram_model_memory=lowvram_model_memory, *args, **kwargs)
if lowvram_model_memory > 0:
self._patch_lowvram_extras(device_to=device_to)
return to_return
def _patch_lowvram_extras(self, device_to=None):
if self.model.motion_wrapper is not None:
# figure out the tensors (likely pe's) that should be cast to device besides just the named_modules
remaining_tensors = list(self.model.motion_wrapper.state_dict().keys())
@@ -134,6 +138,10 @@ class SparseModelPatcher(ModelPatcher):
if device_to is not None:
comfy.utils.set_attr(self.model.motion_wrapper, key, comfy.utils.get_attr(self.model.motion_wrapper, key).to(device_to))
# NOTE: no longer called by ComfyUI, but here for backwards compatibility
def patch_model_lowvram(self, device_to=None, *args, **kwargs):
patched_model = super().patch_model_lowvram(device_to, *args, **kwargs)
self._patch_lowvram_extras(device_to=device_to)
return patched_model
def clone(self):