From 74320a78e3848416b7a7d959850c8b40fe80f3d7 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 12 Nov 2024 18:14:21 -0600 Subject: [PATCH] Replaced SparseModelPatcher with native ModelPatcher and callbacks --- adv_control/control.py | 4 +- adv_control/control_sparsectrl.py | 77 +++++++++++-------------------- 2 files changed, 29 insertions(+), 52 deletions(-) diff --git a/adv_control/control.py b/adv_control/control.py index 7318ef3..4706292 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -11,7 +11,7 @@ import comfy.controlnet as comfy_cn from comfy.controlnet import ControlBase, ControlNet, ControlLora, T2IAdapter, StrengthType from comfy.model_patcher import ModelPatcher -from .control_sparsectrl import SparseModelPatcher, SparseControlNet, SparseCtrlMotionWrapper, SparseSettings, SparseConst +from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseSettings, SparseConst, create_sparse_modelpatcher from .control_lllite import LLLiteModule, LLLitePatch, load_controllllite from .control_svd import svd_unet_config_from_diffusers_unet, SVDControlNet, svd_unet_to_diffusers from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, AbstractPreprocWrapper, ControlWeightType, ControlWeights, WeightTypeException, @@ -313,7 +313,7 @@ class SVDControlNetAdvanced(ControlNetAdvanced): class SparseCtrlAdvanced(ControlNetAdvanced): def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, sparse_settings: SparseSettings=None, global_average_pooling=False, load_device=None, manual_cast_dtype=None): super().__init__(control_model=control_model, timestep_keyframes=timestep_keyframes, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype) - self.control_model_wrapped = SparseModelPatcher(self.control_model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device()) + self.control_model_wrapped = create_sparse_modelpatcher(self.control_model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device()) self.add_compatible_weight(ControlWeightType.SPARSECTRL) self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints if self.control_model.use_simplified_conditioning_embedding: diff --git a/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index 145abc7..214b389 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -25,6 +25,7 @@ from comfy.ldm.modules.attention import attention_basic, attention_pytorch, atte from comfy.ldm.modules.attention import FeedForward, SpatialTransformer from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequential from comfy.model_patcher import ModelPatcher +from comfy.patcher_extension import CallbacksMP import comfy.ops import comfy.model_management import comfy.utils @@ -110,24 +111,22 @@ class SparseControlNet(ControlNetCLDM): return {"middle": out_middle, "output": out_output} -class SparseModelPatcher(ModelPatcher): - def __init__(self, *args, **kwargs): - self.model: SparseControlNet - super().__init__(*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) - self._handle_float8_pe_tensors() - return to_return +def create_sparse_modelpatcher(model, load_device, offload_device): + patcher = ModelPatcher(model, load_device=load_device, offload_device=offload_device) + acn = "ACN" + patcher.add_callback_with_key(CallbacksMP.ON_LOAD, acn, _patch_lowvram_extras) + patcher.add_callback_with_key(CallbacksMP.ON_LOAD, acn, _handle_float8_pe_tensors) + return patcher - def _patch_lowvram_extras(self, device_to=None): - if self.model.motion_wrapper is not None: + +def _patch_lowvram_extras(self: ModelPatcher, device_to, lowvram_model_memory, force_patch_weights, full_load, *args, **kwargs): + if lowvram_model_memory > 0: + motion_wrapper: SparseCtrlMotionWrapper = self.model.motion_wrapper + if 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()) + remaining_tensors = list(motion_wrapper.state_dict().keys()) named_modules = [] - for n, _ in self.model.motion_wrapper.named_modules(): + for n, _ in motion_wrapper.named_modules(): named_modules.append(n) named_modules.append(f"{n}.weight") named_modules.append(f"{n}.bias") @@ -138,43 +137,21 @@ class SparseModelPatcher(ModelPatcher): for key in remaining_tensors: self.patch_weight_to_device(key, device_to) 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)) + comfy.utils.set_attr(motion_wrapper, key, comfy.utils.get_attr(motion_wrapper, key).to(device_to)) - def _handle_float8_pe_tensors(self): - if self.model.motion_wrapper is not None: - remaining_tensors = list(self.model.motion_wrapper.state_dict().keys()) - pe_tensors = [x for x in remaining_tensors if '.pe' in x] - is_first = True - for key in pe_tensors: - if is_first: - is_first = False - if comfy.utils.get_attr(self.model.motion_wrapper, key).dtype not in [torch.float8_e5m2, torch.float8_e4m3fn]: - break - comfy.utils.set_attr(self.model.motion_wrapper, key, comfy.utils.get_attr(self.model.motion_wrapper, key).half()) - # 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): - # normal ModelPatcher clone actions - n = SparseModelPatcher(self.model, self.load_device, self.offload_device, self.size, weight_inplace_update=self.weight_inplace_update) - n.patches = {} - for k in self.patches: - n.patches[k] = self.patches[k][:] - if hasattr(n, "patches_uuid"): - self.patches_uuid = n.patches_uuid - - n.object_patches = self.object_patches.copy() - n.model_options = copy.deepcopy(self.model_options) - if hasattr(n, "model_keys"): - n.model_keys = self.model_keys - if hasattr(n, "backup"): - self.backup = n.backup - if hasattr(n, "object_patches_backup"): - self.object_patches_backup = n.object_patches_backup +def _handle_float8_pe_tensors(self: ModelPatcher, *args, **kwargs): + motion_wrapper: SparseCtrlMotionWrapper = self.model.motion_wrapper + if motion_wrapper is not None: + remaining_tensors = list(motion_wrapper.state_dict().keys()) + pe_tensors = [x for x in remaining_tensors if '.pe' in x] + is_first = True + for key in pe_tensors: + if is_first: + is_first = False + if comfy.utils.get_attr(motion_wrapper, key).dtype not in [torch.float8_e5m2, torch.float8_e4m3fn]: + break + comfy.utils.set_attr(motion_wrapper, key, comfy.utils.get_attr(motion_wrapper, key).half()) class PreprocSparseRGBWrapper(AbstractPreprocWrapper):