Refactored patching code to work with latest ComfyUI - NOT backwards compatible with previous ComfyUI versions, sadly

This commit is contained in:
Jedrzej Kosinski
2024-08-20 01:26:47 -05:00
parent d8af8fe2db
commit ff07c13452
+64 -55
View File
@@ -220,53 +220,54 @@ class ModelPatcherAndInjector(ModelPatcher):
combined_patches[key] = current_patches
return combined_patches
def model_patches_to(self, device):
super().model_patches_to(device)
def patch_model(self, device_to=None, patch_weights=True):
def patch_model(self, *args, **kwargs):
was_injected = False
if self.currently_injected:
self.eject_model()
was_injected = True
# first, perform model patching
if patch_weights: # TODO: keep only 'else' portion when don't need to worry about past comfy versions
patched_model = super().patch_model(device_to)
else:
patched_model = super().patch_model(device_to, patch_weights)
# finally, perform motion model injection
self.inject_model()
patched_model = super().patch_model(*args, **kwargs)
# bring injection back to original state
if was_injected and not self.currently_injected:
self.inject_model()
return patched_model
def patch_model_lowvram(self, *args, **kwargs):
def load(self, device_to=None, lowvram_model_memory=0, *args, **kwargs):
self.eject_model()
try:
return super().patch_model_lowvram(*args, **kwargs)
return super().load(device_to=device_to, lowvram_model_memory=lowvram_model_memory, *args, **kwargs)
finally:
# check if any modules have weight_function or bias_function that is not None
# NOTE: this serves no purpose currently, but I have it here for future reasons
for n, m in self.model.named_modules():
if not hasattr(m, "comfy_cast_weights"):
continue
if getattr(m, "weight_function", None) is not None:
self.model_params_lowvram = True
self.model_params_lowvram_keys[f"{n}.weight"] = n
if getattr(m, "bias_function", None) is not None:
self.model_params_lowvram = True
self.model_params_lowvram_keys[f"{n}.bias"] = n
self.inject_model()
if lowvram_model_memory > 0:
self._patch_lowvram_extras()
def _patch_lowvram_extras(self):
# check if any modules have weight_function or bias_function that is not None
# NOTE: this serves no purpose currently, but I have it here for future reasons
self.model_params_lowvram = False
self.model_params_lowvram_keys.clear()
for n, m in self.model.named_modules():
if not hasattr(m, "comfy_cast_weights"):
continue
if getattr(m, "weight_function", None) is not None:
self.model_params_lowvram = True
self.model_params_lowvram_keys[f"{n}.weight"] = n
if getattr(m, "bias_function", None) is not None:
self.model_params_lowvram = True
self.model_params_lowvram_keys[f"{n}.bias"] = n
def unpatch_model(self, device_to=None, unpatch_weights=True):
# first, eject motion model from unet
self.eject_model()
# finally, do normal model unpatching
if unpatch_weights: # TODO: keep only 'else' portion when don't need to worry about past comfy versions
if unpatch_weights:
# handle hooked_patches first
self.clean_hooks()
try:
return super().unpatch_model(device_to)
finally:
self.model_params_lowvram = False
self.model_params_lowvram_keys.clear()
else:
try:
return super().unpatch_model(device_to, unpatch_weights)
finally:
self.model_params_lowvram = False
self.model_params_lowvram_keys.clear()
try:
return super().unpatch_model(device_to, unpatch_weights)
finally:
self.model_params_lowvram = False
self.model_params_lowvram_keys.clear()
def partially_load(self, *args, **kwargs):
# partially_load calls patch_model, but we don't want to inject model in the intermediate call;
@@ -625,7 +626,7 @@ class ModelPatcherCLIPHooks(ModelPatcher):
else:
comfy.utils.set_attr_param(self.model, key, out_weight)
def patch_model(self, device_to=None, patch_weights=True, *args, **kwargs):
def patch_model(self, device_to=None, *args, **kwargs):
if self.desired_lora_hooks is not None:
self.patches_backup = self.patches.copy()
relevant_patches = self.get_combined_hooked_patches(lora_hooks=self.desired_lora_hooks)
@@ -633,23 +634,29 @@ class ModelPatcherCLIPHooks(ModelPatcher):
self.patches.setdefault(key, [])
self.patches[key].extend(relevant_patches[key])
self.current_lora_hooks = self.desired_lora_hooks
return super().patch_model(device_to, patch_weights, *args, **kwargs)
return super().patch_model(device_to, *args, **kwargs)
def patch_model_lowvram(self, *args, **kwargs):
def load(self, device_to=None, lowvram_model_memory=0, *args, **kwargs):
try:
return super().patch_model_lowvram(*args, **kwargs)
return super().load(device_to=device_to, lowvram_model_memory=lowvram_model_memory, *args, **kwargs)
finally:
# check if any modules have weight_function or bias_function that is not None
# NOTE: this serves no purpose currently, but I have it here for future reasons
for n, m in self.model.named_modules():
if not hasattr(m, "comfy_cast_weights"):
continue
if getattr(m, "weight_function", None) is not None:
self.model_params_lowvram = True
self.model_params_lowvram_keys[f"{n}.weight"] = n
if getattr(m, "bias_function", None) is not None:
self.model_params_lowvram = True
self.model_params_lowvram_keys[f"{n}.weight"] = n
if lowvram_model_memory > 0:
self._patch_lowvram_extras()
def _patch_lowvram_extras(self):
# check if any modules have weight_function or bias_function that is not None
# NOTE: this serves no purpose currently, but I have it here for future reasons
self.model_params_lowvram = False
self.model_params_lowvram_keys.clear()
for n, m in self.model.named_modules():
if not hasattr(m, "comfy_cast_weights"):
continue
if getattr(m, "weight_function", None) is not None:
self.model_params_lowvram = True
self.model_params_lowvram_keys[f"{n}.weight"] = n
if getattr(m, "bias_function", None) is not None:
self.model_params_lowvram = True
self.model_params_lowvram_keys[f"{n}.weight"] = n
def unpatch_model(self, device_to=None, unpatch_weights=True, *args, **kwargs):
try:
@@ -797,10 +804,14 @@ class MotionModelPatcher(ModelPatcher):
self.was_within_range = False
self.prev_sub_idxs = None
self.prev_batched_number = None
def patch_model_lowvram(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, *args, **kwargs):
patched_model = super().patch_model_lowvram(device_to, lowvram_model_memory, force_patch_weights, *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):
# figure out the tensors (likely pe's) that should be cast to device besides just the named_modules
remaining_tensors = list(self.model.state_dict().keys())
named_modules = []
@@ -817,8 +828,6 @@ class MotionModelPatcher(ModelPatcher):
if device_to is not None:
comfy.utils.set_attr(self.model, key, comfy.utils.get_attr(self.model, key).to(device_to))
return patched_model
def pre_run(self, model: ModelPatcherAndInjector):
self.cleanup()
self.model.set_scale(self.scale_multival, self.per_block_list)