diff --git a/comfy/controlnet.py b/comfy/controlnet.py index b80634913..3d0e25e89 100644 --- a/comfy/controlnet.py +++ b/comfy/controlnet.py @@ -35,6 +35,9 @@ import comfy.ldm.cascade.controlnet import comfy.cldm.mmdit import comfy.ldm.hydit.controlnet import comfy.ldm.flux.controlnet +from typing import TYPE_CHECKING +if TYPE_CHECKING: + from comfy.hooks import HookGroup def broadcast_image_to(tensor, target_batch_size, batched_number): @@ -78,6 +81,7 @@ class ControlBase: self.concat_mask = False self.extra_concat_orig = [] self.extra_concat = None + self.extra_hooks: HookGroup = None def set_cond_hint(self, cond_hint, strength=1.0, timestep_percent_range=(0.0, 1.0), vae=None, extra_concat=[]): self.cond_hint_original = cond_hint @@ -129,6 +133,7 @@ class ControlBase: c.strength_type = self.strength_type c.concat_mask = self.concat_mask c.extra_concat_orig = self.extra_concat_orig.copy() + c.extra_hooks = self.extra_hooks.clone() if self.extra_hooks else None def inference_memory_requirements(self, dtype): if self.previous_controlnet is not None: diff --git a/comfy/model_base.py b/comfy/model_base.py index a2aa4e317..0857dfe65 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -41,7 +41,7 @@ import comfy.latent_formats import math from typing import TYPE_CHECKING if TYPE_CHECKING: - from model_patcher import ModelPatcher + from comfy.model_patcher import ModelPatcher class ModelType(Enum): EPS = 1 diff --git a/comfy/patcher_extension.py b/comfy/patcher_extension.py index 2d93b2b08..096f556bb 100644 --- a/comfy/patcher_extension.py +++ b/comfy/patcher_extension.py @@ -88,6 +88,8 @@ def get_all_wrappers(wrapper_type: str, transformer_options: dict, is_model_opti class WrapperExecutor: """Handles call stack of wrappers around a function in an ordered manner.""" def __init__(self, original: Callable, class_obj: object, wrappers: list[Callable], idx: int): + # NOTE: class_obj exists so that wrappers surrounding a class method can access + # the class instance at runtime via executor.class_obj self.original = original self.class_obj = class_obj self.wrappers = wrappers.copy() @@ -95,7 +97,7 @@ class WrapperExecutor: self.is_last = idx == len(wrappers) def __call__(self, *args, **kwargs): - """Calls the next wrapper in line or original function, whichever is appropriate.""" + """Calls the next wrapper or original function, whichever is appropriate.""" new_executor = self._create_next_executor() return new_executor.execute(*args, **kwargs) diff --git a/comfy/sampler_helpers.py b/comfy/sampler_helpers.py index a095367ad..f8622b466 100644 --- a/comfy/sampler_helpers.py +++ b/comfy/sampler_helpers.py @@ -30,12 +30,34 @@ def get_models_from_cond(cond, model_type): return models def get_hooks_from_cond(cond, hooks_dict: dict[comfy.hooks.EnumHookType, dict[comfy.hooks.Hook, None]]): + # get hooks from conds, and collect cnets so they can be checked for extra_hooks + cnets: list[ControlBase] = [] for c in cond: if 'hooks' in c: for hook in c['hooks'].hooks: hook: comfy.hooks.Hook with_type = hooks_dict.setdefault(hook.hook_type, {}) with_type[hook] = None + if 'control' in c: + cnets.append(c['control']) + + def get_extra_hooks_from_cnet(cnet: ControlBase, _list: list): + if cnet.extra_hooks is not None: + _list.append(cnet.extra_hooks) + if cnet.previous_controlnet is None: + return _list + return get_extra_hooks_from_cnet(cnet.previous_controlnet, _list) + + hooks_list = [] + cnets = set(cnets) + for base_cnet in cnets: + get_extra_hooks_from_cnet(base_cnet, hooks_list) + extra_hooks = comfy.hooks.HookGroup.combine_all_hooks(hooks_list) + if extra_hooks is not None: + for hook in extra_hooks.hooks: + with_type = hooks_dict.setdefault(hook.hook_type, {}) + with_type[hook] = None + return hooks_dict def convert_cond(cond): diff --git a/comfy/samplers.py b/comfy/samplers.py index 3d8ca1d21..5ae56ce85 100644 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -9,6 +9,7 @@ import collections from comfy import model_management import math import logging +import comfy.samplers import comfy.sampler_helpers import comfy.model_patcher import comfy.patcher_extension @@ -77,6 +78,7 @@ def get_area_and_mult(conds, x_in, timestep_in): for c in model_conds: conditioning[c] = model_conds[c].process_cond(batch_size=x_in.shape[0], device=x_in.device, area=area) + hooks = conds.get('hooks', None) control = conds.get('control', None) patches = None @@ -92,8 +94,8 @@ def get_area_and_mult(conds, x_in, timestep_in): patches['middle_patch'] = [gligen_patch] - cond_obj = collections.namedtuple('cond_obj', ['input_x', 'mult', 'conditioning', 'area', 'control', 'patches', 'uuid']) - return cond_obj(input_x, mult, conditioning, area, control, patches, conds['uuid']) + cond_obj = collections.namedtuple('cond_obj', ['input_x', 'mult', 'conditioning', 'area', 'control', 'patches', 'uuid', 'hooks']) + return cond_obj(input_x, mult, conditioning, area, control, patches, conds['uuid'], hooks) def cond_equal_size(c1, c2): if c1 is c2: @@ -179,20 +181,19 @@ def finalize_default_conds(model: 'BaseModel', hooked_to_run: dict[comfy.hooks.H continue # replace p's mult with calculated mult p = p._replace(mult=mult) - hooks: comfy.hooks.HookGroup = x.get('hooks', None) - if hooks is not None: - model.current_patcher.prepare_hook_patches_current_keyframe(timestep, hooks) - hooked_to_run.setdefault(hooks, list()) - hooked_to_run[hooks] += [(p, i)] + if p.hooks is not None: + model.current_patcher.prepare_hook_patches_current_keyframe(timestep, p.hooks) + hooked_to_run.setdefault(p.hooks, list()) + hooked_to_run[p.hooks] += [(p, i)] def calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: torch.Tensor, timestep, model_options): executor = comfy.patcher_extension.WrapperExecutor.new_executor( - outer_calc_cond_batch, + _calc_cond_batch, comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.CALC_COND_BATCH, model_options, is_model_options=True) ) return executor.execute(model, conds, x_in, timestep, model_options) -def outer_calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: torch.Tensor, timestep, model_options): +def _calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: torch.Tensor, timestep, model_options): out_conds = [] out_counts = [] # separate conds by matching hooks @@ -215,11 +216,10 @@ def outer_calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: tor p = comfy.samplers.get_area_and_mult(x, x_in, timestep) if p is None: continue - hooks: comfy.hooks.HookGroup = x.get('hooks', None) - if hooks is not None: - model.current_patcher.prepare_hook_patches_current_keyframe(timestep, hooks) - hooked_to_run.setdefault(hooks, list()) - hooked_to_run[hooks] += [(p, i)] + if p.hooks is not None: + model.current_patcher.prepare_hook_patches_current_keyframe(timestep, p.hooks) + hooked_to_run.setdefault(p.hooks, list()) + hooked_to_run[p.hooks] += [(p, i)] default_conds.append(default_c) if has_default_conds: