diff --git a/comfy/ldm/modules/diffusionmodules/openaimodel.py b/comfy/ldm/modules/diffusionmodules/openaimodel.py index 6454be245..3f7fee708 100644 --- a/comfy/ldm/modules/diffusionmodules/openaimodel.py +++ b/comfy/ldm/modules/diffusionmodules/openaimodel.py @@ -15,6 +15,7 @@ from .util import ( ) from ..attention import SpatialTransformer, SpatialVideoTransformer, default from comfy.ldm.util import exists +import comfy.patcher_extension import comfy.ops ops = comfy.ops.disable_weight_init @@ -828,6 +829,13 @@ class UNetModel(nn.Module): ) def forward(self, x, timesteps=None, context=None, y=None, control=None, transformer_options={}, **kwargs): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._forward, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options) + ).execute(x, timesteps, context, y, control, transformer_options, **kwargs) + + def _forward(self, x, timesteps=None, context=None, y=None, control=None, transformer_options={}, **kwargs): """ Apply the model to an input batch. :param x: an [N x C x ...] Tensor of inputs. diff --git a/comfy/model_base.py b/comfy/model_base.py index 2ccbd0ff5..a2aa4e317 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -32,6 +32,7 @@ import comfy.ldm.audio.embedders import comfy.ldm.flux.model import comfy.model_management +import comfy.patcher_extension import comfy.conds import comfy.ops from enum import Enum @@ -123,6 +124,13 @@ class BaseModel(torch.nn.Module): self.memory_usage_factor = model_config.memory_usage_factor def apply_model(self, x, t, c_concat=None, c_crossattn=None, control=None, transformer_options={}, **kwargs): + return comfy.patcher_extension.WrapperExecutor.new_class_executor( + self._apply_model, + self, + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.APPLY_MODEL, transformer_options) + ).execute(x, t, c_concat, c_crossattn, control, transformer_options, **kwargs) + + def _apply_model(self, x, t, c_concat=None, c_crossattn=None, control=None, transformer_options={}, **kwargs): sigma = t xc = self.model_sampling.calculate_input(sigma, x) if c_concat is not None: diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index f6e5bab83..21c758858 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -31,6 +31,7 @@ import comfy.float import comfy.model_management import comfy.lora import comfy.hooks +import comfy.patcher_extension from comfy.patcher_extension import CallbacksMP, WrappersMP, PatcherInjection from comfy.comfy_types import UnetWrapperFunction @@ -81,15 +82,7 @@ def set_model_options_pre_cfg_function(model_options, pre_cfg_function, disable_ return model_options def create_model_options_clone(orig_model_options: dict): - def copy_nested_dicts(input_dict: dict): - new_dict = input_dict.copy() - for key, value in input_dict.items(): - if isinstance(value, dict): - new_dict[key] = copy_nested_dicts(value) - elif isinstance(value, list): - new_dict[key] = value.copy() - return new_dict - return copy_nested_dicts(orig_model_options) + return comfy.patcher_extension.copy_nested_dicts(orig_model_options) def create_hook_patches_clone(orig_hook_patches): new_hook_patches = {} diff --git a/comfy/patcher_extension.py b/comfy/patcher_extension.py index 980fb8eb6..2d93b2b08 100644 --- a/comfy/patcher_extension.py +++ b/comfy/patcher_extension.py @@ -26,6 +26,27 @@ class CallbacksMP: cls.ON_EJECT_MODEL: {None: []}, } +def add_callback(call_type: str, callback: Callable, transformer_options: dict, is_model_options=False): + add_callback_with_key(call_type, None, callback, transformer_options, is_model_options) + +def add_callback_with_key(call_type: str, key: str, callback: Callable, transformer_options: dict, is_model_options=False): + if is_model_options: + transformer_options = transformer_options.get("transformer_options", {}) + callbacks: dict[str, dict[str, list]] = transformer_options.get("callbacks", {}) + if call_type not in callbacks: + raise Exception(f"Callback '{call_type}' is not recognized.") + c = callbacks[call_type].setdefault(key, []) + c.append(callback) + +def get_all_callbacks(call_type: str, transformer_options: dict, is_model_options=False): + if is_model_options: + transformer_options = transformer_options.get("transformer_options", {}) + c_list = [] + callbacks: dict[str, list] = transformer_options.get("callbacks", {}) + for c in callbacks.get(call_type, {}).values(): + c_list.extend(c) + return c_list + class WrappersMP: OUTER_SAMPLE = "outer_sample" SAMPLER_SAMPLE = "sampler_sample" @@ -43,6 +64,27 @@ class WrappersMP: cls.DIFFUSION_MODEL: {None: []}, } +def add_wrapper(wrapper_type: str, wrapper: Callable, transformer_options: dict, is_model_options=False): + add_wrapper_with_key(wrapper_type, None, wrapper, transformer_options, is_model_options) + +def add_wrapper_with_key(wrapper_type: str, key: str, wrapper: Callable, transformer_options: dict, is_model_options=False): + if is_model_options: + transformer_options = transformer_options.get("transformer_options", {}) + wrappers: dict[str, dict[str, list]] = transformer_options.get("wrappers", {}) + if wrapper_type not in wrappers: + raise Exception(f"Wrapper '{wrapper_type}' is not recognized.") + w = wrappers[wrapper_type].setdefault(key, []) + w.append(wrapper) + +def get_all_wrappers(wrapper_type: str, transformer_options: dict, is_model_options=False): + if is_model_options: + transformer_options = transformer_options.get("transformer_options", {}) + w_list = [] + wrappers: dict[str, list] = transformer_options.get("wrappers", {}) + for w in wrappers.get(wrapper_type, {}).values(): + w_list.extend(w) + return w_list + 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): @@ -85,3 +127,24 @@ class PatcherInjection: def __init__(self, inject: Callable, eject: Callable): self.inject = inject self.eject = eject + +def copy_nested_dicts(input_dict: dict): + new_dict = input_dict.copy() + for key, value in input_dict.items(): + if isinstance(value, dict): + new_dict[key] = copy_nested_dicts(value) + elif isinstance(value, list): + new_dict[key] = value.copy() + return new_dict + +def merge_nested_dicts(dict1: dict, dict2: dict): + merged_dict = copy_nested_dicts(dict1) + for key, value in dict2.items(): + if isinstance(value, dict): + curr_value = merged_dict.setdefault(key, {}) + merged_dict[key] = merge_nested_dicts(value, curr_value) + elif isinstance(value, list): + merged_dict.setdefault(key, []).extend(value) + else: + merged_dict[key] = value + return merged_dict diff --git a/comfy/sampler_helpers.py b/comfy/sampler_helpers.py index e1c0d9db9..a095367ad 100644 --- a/comfy/sampler_helpers.py +++ b/comfy/sampler_helpers.py @@ -4,6 +4,7 @@ import torch import comfy.model_management import comfy.conds import comfy.hooks +import comfy.patcher_extension from typing import TYPE_CHECKING if TYPE_CHECKING: from comfy.model_patcher import ModelPatcher @@ -105,9 +106,13 @@ def cleanup_models(conds, models): cleanup_additional_models(set(control_cleanup)) -def prepare_model_patcher(model: 'ModelPatcher', conds): +def prepare_model_patcher(model: 'ModelPatcher', conds, model_options: dict): # check for hooks in conds - if not registered, see if can be applied hooks = {} for k in conds: get_hooks_from_cond(conds[k], hooks) model.register_all_hook_patches(hooks, comfy.hooks.EnumWeightTarget.Model) + # add wrappers and callbacks from ModelPatcher to transformer_options + model_options["transformer_options"]["wrappers"] = comfy.patcher_extension.copy_nested_dicts(model.wrappers) + model_options["transformer_options"]["callbacks"] = comfy.patcher_extension.copy_nested_dicts(model.callbacks) + # TODO: add wrappers and callbacks from registered hooks for functions called prior to calc_batch_conds diff --git a/comfy/samplers.py b/comfy/samplers.py index 943b7cbb9..3d8ca1d21 100644 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -188,7 +188,7 @@ def finalize_default_conds(model: 'BaseModel', hooked_to_run: dict[comfy.hooks.H 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, - model.current_patcher.get_all_wrappers(comfy.patcher_extension.WrappersMP.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) @@ -808,7 +808,7 @@ class CFGGuider: executor = comfy.patcher_extension.WrapperExecutor.new_class_executor( sampler.sample, sampler, - self.model_patcher.get_all_wrappers(comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE) + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE, extra_args["model_options"], is_model_options=True) ) samples = executor.execute(self, sigmas, extra_args, callback, noise, latent_image, denoise_mask, disable_pbar) return self.inner_model.process_latent_out(samples.to(torch.float32)) @@ -844,14 +844,17 @@ class CFGGuider: self.conds[k] = list(map(lambda a: a.copy(), self.original_conds[k])) try: - comfy.sampler_helpers.prepare_model_patcher(self.model_patcher, self.conds) + orig_model_options = self.model_options + self.model_options = comfy.model_patcher.create_model_options_clone(self.model_options) + comfy.sampler_helpers.prepare_model_patcher(self.model_patcher, self.conds, self.model_options) executor = comfy.patcher_extension.WrapperExecutor.new_class_executor( self.outer_sample, self, - self.model_patcher.get_all_wrappers(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE) + comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, self.model_options, is_model_options=True) ) output = executor.execute(noise, latent_image, sampler, sigmas, denoise_mask, callback, disable_pbar, seed) finally: + self.model_options = orig_model_options self.model_patcher.restore_hook_patches() del self.conds