mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-10-10 00:17:12 +08:00
Started scaffolding for other hook types, refactored get_hooks_from_cond to organize hooks by type
This commit is contained in:
parent
59d72b4050
commit
5f450d3351
101
comfy/hooks.py
101
comfy/hooks.py
@ -5,7 +5,7 @@ import torch
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from comfy.model_patcher import ModelPatcher
|
from comfy.model_patcher import ModelPatcher, PatcherInjection
|
||||||
from comfy.model_base import BaseModel
|
from comfy.model_base import BaseModel
|
||||||
from comfy.sd import CLIP
|
from comfy.sd import CLIP
|
||||||
import comfy.lora
|
import comfy.lora
|
||||||
@ -19,7 +19,11 @@ class EnumHookMode(enum.Enum):
|
|||||||
class EnumHookType(enum.Enum):
|
class EnumHookType(enum.Enum):
|
||||||
Weight = "weight"
|
Weight = "weight"
|
||||||
Patch = "patch"
|
Patch = "patch"
|
||||||
AddModel = "addmodel"
|
ObjectPatch = "object_patch"
|
||||||
|
AddModels = "add_models"
|
||||||
|
AddCallback = "add_callback"
|
||||||
|
SetInjections = "add_injections"
|
||||||
|
AddWrapper = "add_wrapper"
|
||||||
|
|
||||||
class EnumWeightTarget(enum.Enum):
|
class EnumWeightTarget(enum.Enum):
|
||||||
Model = "model"
|
Model = "model"
|
||||||
@ -126,18 +130,94 @@ class PatchHook(Hook):
|
|||||||
c.patches = self.patches
|
c.patches = self.patches
|
||||||
return c
|
return c
|
||||||
|
|
||||||
class AddModelHook(Hook):
|
def add_hook_patches(self, model: 'ModelPatcher'):
|
||||||
def __init__(self, model: 'ModelPatcher'):
|
pass
|
||||||
super().__init__(hook_type=EnumHookType.AddModel)
|
|
||||||
self.model = model
|
class ObjectPatchHook(Hook):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__(hook_type=EnumHookType.ObjectPatch)
|
||||||
|
self.object_patches: Dict = None
|
||||||
|
|
||||||
def clone(self, subtype: Callable=None):
|
def clone(self, subtype: Callable=None):
|
||||||
if subtype is None:
|
if subtype is None:
|
||||||
subtype = type(self)
|
subtype = type(self)
|
||||||
c: AddModelHook = super().clone(subtype)
|
c: ObjectPatchHook = super().clone(subtype)
|
||||||
c.model = self.model
|
c.object_patches = self.object_patches
|
||||||
return c
|
return c
|
||||||
|
|
||||||
|
def add_hook_object_patches(self, model: 'ModelPatcher'):
|
||||||
|
pass
|
||||||
|
|
||||||
|
class AddModelsHook(Hook):
|
||||||
|
def __init__(self, key: str=None, models: List['ModelPatcher']=None):
|
||||||
|
super().__init__(hook_type=EnumHookType.AddModels)
|
||||||
|
self.key = key
|
||||||
|
self.models = models
|
||||||
|
self.append_when_same = True
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: AddModelsHook = super().clone(subtype)
|
||||||
|
c.key = self.key
|
||||||
|
c.models = self.models.copy() if self.models else self.models
|
||||||
|
c.append_when_same = self.append_when_same
|
||||||
|
return c
|
||||||
|
|
||||||
|
def add_hook_models(self, model: 'ModelPatcher'):
|
||||||
|
pass
|
||||||
|
|
||||||
|
class AddCallbackHook(Hook):
|
||||||
|
def __init__(self, key: str=None, callback: Callable=None):
|
||||||
|
super().__init__(hook_type=EnumHookType.AddCallback)
|
||||||
|
self.key = key
|
||||||
|
self.callback = callback
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: AddCallbackHook = super().clone(subtype)
|
||||||
|
c.key = self.key
|
||||||
|
c.callback = self.callback
|
||||||
|
return c
|
||||||
|
|
||||||
|
def add_hook_callback(self, model: 'ModelPatcher'):
|
||||||
|
pass
|
||||||
|
|
||||||
|
class SetInjectionsHook(Hook):
|
||||||
|
def __init__(self, key: str=None, injections: List['PatcherInjection']=None):
|
||||||
|
super().__init__(hook_type=EnumHookType.SetInjections)
|
||||||
|
self.key = key
|
||||||
|
self.injections = injections
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: SetInjectionsHook = super().clone(subtype)
|
||||||
|
c.key = self.key
|
||||||
|
c.injections = self.injections.copy() if self.injections else self.injections
|
||||||
|
return c
|
||||||
|
|
||||||
|
def add_hook_injections(self, model: 'ModelPatcher'):
|
||||||
|
pass
|
||||||
|
|
||||||
|
class AddWrapperHook(Hook):
|
||||||
|
def __init__(self, key: str=None, wrapper: Callable=None):
|
||||||
|
super().__init__(hook_type=EnumHookType.AddWrapper)
|
||||||
|
self.key = key
|
||||||
|
self.wrapper = wrapper
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: AddWrapperHook = super().clone(subtype)
|
||||||
|
c.key = self.key
|
||||||
|
c.wrapper = self.wrapper
|
||||||
|
return c
|
||||||
|
|
||||||
|
def add_hook_wrapper(self, model: 'ModelPatcher'):
|
||||||
|
pass
|
||||||
|
|
||||||
class HookGroup:
|
class HookGroup:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.hooks: List[Hook] = []
|
self.hooks: List[Hook] = []
|
||||||
@ -167,9 +247,10 @@ class HookGroup:
|
|||||||
hook.hook_keyframe = hook_kf
|
hook.hook_keyframe = hook_kf
|
||||||
|
|
||||||
def get_dict_repr(self):
|
def get_dict_repr(self):
|
||||||
d = {}
|
d: Dict[EnumHookType, Dict[Hook, None]] = {}
|
||||||
for hook in self.hooks:
|
for hook in self.hooks:
|
||||||
d[hook] = None
|
with_type = d.setdefault(hook.hook_type, {})
|
||||||
|
with_type[hook] = None
|
||||||
return d
|
return d
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@ -809,13 +809,12 @@ class ModelPatcher:
|
|||||||
if cached_group.contains(hook):
|
if cached_group.contains(hook):
|
||||||
self.cached_hook_patches.pop(cached_group)
|
self.cached_hook_patches.pop(cached_group)
|
||||||
|
|
||||||
def register_all_hook_patches(self, hooks_dict: Dict[comfy.hooks.Hook, None], target: comfy.hooks.EnumWeightTarget):
|
def register_all_hook_patches(self, hooks_dict: Dict[comfy.hooks.EnumHookType, Dict[comfy.hooks.Hook, None]], target: comfy.hooks.EnumWeightTarget):
|
||||||
self.restore_hook_patches()
|
self.restore_hook_patches()
|
||||||
weight_hooks_to_register: List[comfy.hooks.WeightHook] = []
|
weight_hooks_to_register: List[comfy.hooks.WeightHook] = []
|
||||||
for hook in hooks_dict:
|
for hook in hooks_dict.get(comfy.hooks.EnumHookType.Weight, {}):
|
||||||
if hook.hook_type == comfy.hooks.EnumHookType.Weight:
|
if hook.hook_ref not in self.hook_patches:
|
||||||
if hook.hook_ref not in self.hook_patches:
|
weight_hooks_to_register.append(hook)
|
||||||
weight_hooks_to_register.append(hook)
|
|
||||||
if len(weight_hooks_to_register) > 0:
|
if len(weight_hooks_to_register) > 0:
|
||||||
self.hook_patches_backup = create_hook_patches_clone(self.hook_patches)
|
self.hook_patches_backup = create_hook_patches_clone(self.hook_patches)
|
||||||
for hook in weight_hooks_to_register:
|
for hook in weight_hooks_to_register:
|
||||||
|
|||||||
@ -26,15 +26,14 @@ def get_models_from_cond(cond, model_type):
|
|||||||
models += [c[model_type]]
|
models += [c[model_type]]
|
||||||
return models
|
return models
|
||||||
|
|
||||||
def get_hooks_from_cond(cond, filter_types: List[comfy.hooks.EnumHookType]=None):
|
def get_hooks_from_cond(cond, hooks_dict: Dict[comfy.hooks.EnumHookType, Dict[comfy.hooks.Hook, None]]):
|
||||||
hooks: Dict[comfy.hooks.Hook, None] = {}
|
|
||||||
for c in cond:
|
for c in cond:
|
||||||
if 'hooks' in c:
|
if 'hooks' in c:
|
||||||
for hook in c['hooks'].hooks:
|
for hook in c['hooks'].hooks:
|
||||||
hook: comfy.hooks.Hook
|
hook: comfy.hooks.Hook
|
||||||
if not filter_types or hook.hook_type in filter_types:
|
with_type = hooks_dict.setdefault(hook.hook_type, {})
|
||||||
hooks[hook] = None
|
with_type[hook] = None
|
||||||
return hooks
|
return hooks_dict
|
||||||
|
|
||||||
def convert_cond(cond):
|
def convert_cond(cond):
|
||||||
out = []
|
out = []
|
||||||
@ -53,13 +52,13 @@ def get_additional_models(conds, dtype):
|
|||||||
cnets: List[ControlBase] = []
|
cnets: List[ControlBase] = []
|
||||||
gligen = []
|
gligen = []
|
||||||
add_models = []
|
add_models = []
|
||||||
hooks: Dict[comfy.hooks.AddModelHook, None] = {}
|
hooks: Dict[comfy.hooks.EnumHookType, Dict[comfy.hooks.Hook, None]] = {}
|
||||||
|
|
||||||
for k in conds:
|
for k in conds:
|
||||||
cnets += get_models_from_cond(conds[k], "control")
|
cnets += get_models_from_cond(conds[k], "control")
|
||||||
gligen += get_models_from_cond(conds[k], "gligen")
|
gligen += get_models_from_cond(conds[k], "gligen")
|
||||||
add_models += get_models_from_cond(conds[k], "additional_models")
|
add_models += get_models_from_cond(conds[k], "additional_models")
|
||||||
hooks.update(get_hooks_from_cond(conds[k], [comfy.hooks.EnumHookType.AddModel]))
|
get_hooks_from_cond(conds[k], hooks)
|
||||||
|
|
||||||
control_nets = set(cnets)
|
control_nets = set(cnets)
|
||||||
|
|
||||||
@ -70,7 +69,7 @@ def get_additional_models(conds, dtype):
|
|||||||
inference_memory += m.inference_memory_requirements(dtype)
|
inference_memory += m.inference_memory_requirements(dtype)
|
||||||
|
|
||||||
gligen = [x[1] for x in gligen]
|
gligen = [x[1] for x in gligen]
|
||||||
hook_models = [x.model for x in hooks]
|
hook_models = [x.model for x in hooks.get(comfy.hooks.EnumHookType.AddModels, {}).keys()]
|
||||||
models = control_models + gligen + add_models + hook_models
|
models = control_models + gligen + add_models + hook_models
|
||||||
|
|
||||||
return models, inference_memory
|
return models, inference_memory
|
||||||
@ -108,5 +107,5 @@ def prepare_model_patcher(model: 'ModelPatcher', conds):
|
|||||||
# check for hooks in conds - if not registered, see if can be applied
|
# check for hooks in conds - if not registered, see if can be applied
|
||||||
hooks = {}
|
hooks = {}
|
||||||
for k in conds:
|
for k in conds:
|
||||||
hooks.update(get_hooks_from_cond(conds[k]))
|
get_hooks_from_cond(conds[k], hooks)
|
||||||
model.register_all_hook_patches(hooks, comfy.hooks.EnumWeightTarget.Model)
|
model.register_all_hook_patches(hooks, comfy.hooks.EnumWeightTarget.Model)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user