mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-10-08 05:17:05 +08:00
Refactored code to store wrappers and callbacks in transformer_options, added apply_model and diffusion_model.forward wrappers
This commit is contained in:
parent
51e8d5554c
commit
0fbefb8428
@ -15,6 +15,7 @@ from .util import (
|
|||||||
)
|
)
|
||||||
from ..attention import SpatialTransformer, SpatialVideoTransformer, default
|
from ..attention import SpatialTransformer, SpatialVideoTransformer, default
|
||||||
from comfy.ldm.util import exists
|
from comfy.ldm.util import exists
|
||||||
|
import comfy.patcher_extension
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
ops = comfy.ops.disable_weight_init
|
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):
|
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.
|
Apply the model to an input batch.
|
||||||
:param x: an [N x C x ...] Tensor of inputs.
|
:param x: an [N x C x ...] Tensor of inputs.
|
||||||
|
|||||||
@ -32,6 +32,7 @@ import comfy.ldm.audio.embedders
|
|||||||
import comfy.ldm.flux.model
|
import comfy.ldm.flux.model
|
||||||
|
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
|
import comfy.patcher_extension
|
||||||
import comfy.conds
|
import comfy.conds
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
@ -123,6 +124,13 @@ class BaseModel(torch.nn.Module):
|
|||||||
self.memory_usage_factor = model_config.memory_usage_factor
|
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):
|
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
|
sigma = t
|
||||||
xc = self.model_sampling.calculate_input(sigma, x)
|
xc = self.model_sampling.calculate_input(sigma, x)
|
||||||
if c_concat is not None:
|
if c_concat is not None:
|
||||||
|
|||||||
@ -31,6 +31,7 @@ import comfy.float
|
|||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
import comfy.lora
|
import comfy.lora
|
||||||
import comfy.hooks
|
import comfy.hooks
|
||||||
|
import comfy.patcher_extension
|
||||||
from comfy.patcher_extension import CallbacksMP, WrappersMP, PatcherInjection
|
from comfy.patcher_extension import CallbacksMP, WrappersMP, PatcherInjection
|
||||||
from comfy.comfy_types import UnetWrapperFunction
|
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
|
return model_options
|
||||||
|
|
||||||
def create_model_options_clone(orig_model_options: dict):
|
def create_model_options_clone(orig_model_options: dict):
|
||||||
def copy_nested_dicts(input_dict: dict):
|
return comfy.patcher_extension.copy_nested_dicts(orig_model_options)
|
||||||
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)
|
|
||||||
|
|
||||||
def create_hook_patches_clone(orig_hook_patches):
|
def create_hook_patches_clone(orig_hook_patches):
|
||||||
new_hook_patches = {}
|
new_hook_patches = {}
|
||||||
|
|||||||
@ -26,6 +26,27 @@ class CallbacksMP:
|
|||||||
cls.ON_EJECT_MODEL: {None: []},
|
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:
|
class WrappersMP:
|
||||||
OUTER_SAMPLE = "outer_sample"
|
OUTER_SAMPLE = "outer_sample"
|
||||||
SAMPLER_SAMPLE = "sampler_sample"
|
SAMPLER_SAMPLE = "sampler_sample"
|
||||||
@ -43,6 +64,27 @@ class WrappersMP:
|
|||||||
cls.DIFFUSION_MODEL: {None: []},
|
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:
|
class WrapperExecutor:
|
||||||
"""Handles call stack of wrappers around a function in an ordered manner."""
|
"""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):
|
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):
|
def __init__(self, inject: Callable, eject: Callable):
|
||||||
self.inject = inject
|
self.inject = inject
|
||||||
self.eject = eject
|
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
|
||||||
|
|||||||
@ -4,6 +4,7 @@ import torch
|
|||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
import comfy.conds
|
import comfy.conds
|
||||||
import comfy.hooks
|
import comfy.hooks
|
||||||
|
import comfy.patcher_extension
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from comfy.model_patcher import ModelPatcher
|
from comfy.model_patcher import ModelPatcher
|
||||||
@ -105,9 +106,13 @@ def cleanup_models(conds, models):
|
|||||||
|
|
||||||
cleanup_additional_models(set(control_cleanup))
|
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
|
# check for hooks in conds - if not registered, see if can be applied
|
||||||
hooks = {}
|
hooks = {}
|
||||||
for k in conds:
|
for k in conds:
|
||||||
get_hooks_from_cond(conds[k], hooks)
|
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)
|
||||||
|
# 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
|
||||||
|
|||||||
@ -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):
|
def calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: torch.Tensor, timestep, model_options):
|
||||||
executor = comfy.patcher_extension.WrapperExecutor.new_executor(
|
executor = comfy.patcher_extension.WrapperExecutor.new_executor(
|
||||||
outer_calc_cond_batch,
|
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)
|
return executor.execute(model, conds, x_in, timestep, model_options)
|
||||||
|
|
||||||
@ -808,7 +808,7 @@ class CFGGuider:
|
|||||||
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||||
sampler.sample,
|
sampler.sample,
|
||||||
sampler,
|
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)
|
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))
|
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]))
|
self.conds[k] = list(map(lambda a: a.copy(), self.original_conds[k]))
|
||||||
|
|
||||||
try:
|
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(
|
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||||
self.outer_sample,
|
self.outer_sample,
|
||||||
self,
|
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)
|
output = executor.execute(noise, latent_image, sampler, sigmas, denoise_mask, callback, disable_pbar, seed)
|
||||||
finally:
|
finally:
|
||||||
|
self.model_options = orig_model_options
|
||||||
self.model_patcher.restore_hook_patches()
|
self.model_patcher.restore_hook_patches()
|
||||||
|
|
||||||
del self.conds
|
del self.conds
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user