diff --git a/__init__.py b/__init__.py index 5e3e0b2..bb650ed 100644 --- a/__init__.py +++ b/__init__.py @@ -1,9 +1,11 @@ from .adv_control.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS from .adv_control import documentation from .adv_control.dinklink import init_dinklink +from .adv_control.sampling import prepare_dinklink_acn_wrapper WEB_DIRECTORY = "./web" __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', "WEB_DIRECTORY"] documentation.format_descriptions(NODE_CLASS_MAPPINGS) init_dinklink() +prepare_dinklink_acn_wrapper() diff --git a/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index 616aa70..5b5a53d 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -30,6 +30,7 @@ import comfy.ops import comfy.model_management import comfy.utils +from .dinklink import get_AnimateDiffModel, get_AnimateDiffInfo from .logger import logger from .utils import (BIGMAX, AbstractPreprocWrapper, disable_weight_init_clean_groupnorm, prepare_mask_batch, broadcast_image_to_extend, extend_to_batch_size) diff --git a/adv_control/dinklink.py b/adv_control/dinklink.py index a0e8b39..d591d3b 100644 --- a/adv_control/dinklink.py +++ b/adv_control/dinklink.py @@ -12,27 +12,51 @@ #################################################################################################### from __future__ import annotations import comfy.hooks -from comfy.patcher_extension import WrappersMP - -from .sampling import acn_outer_sample_wrapper -from .utils import WrapperConsts DINKLINK = "__DINKLINK" + def init_dinklink(): - if not hasattr(comfy.hooks, DINKLINK): - setattr(comfy.hooks, DINKLINK, {}) + create_dinklink() prepare_dinklink() +def create_dinklink(): + if not hasattr(comfy.hooks, DINKLINK): + setattr(comfy.hooks, DINKLINK, {}) def get_dinklink() -> dict[str, dict[str]]: + create_dinklink() return getattr(comfy.hooks, DINKLINK) + +class DinkLinkConst: + VERSION = "version" + # ADE + ADE = "ADE" + ADE_ANIMATEDIFFMODEL = "AnimateDiffModel" + ADE_ANIMATEDIFFINFO = "AnimateDiffInfo" + def prepare_dinklink(): - # expose acn_sampler_sample_wrapper + pass + +def get_AnimateDiffModel(throw_exception=True): d = get_dinklink() - link_acn = d.setdefault(WrapperConsts.ACN, {}) - link_acn[WrapperConsts.VERSION] = 1 - link_acn[WrapperConsts.ACN_CREATE_SAMPLER_SAMPLE_WRAPPER] = (WrappersMP.OUTER_SAMPLE, - WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY, - acn_outer_sample_wrapper) + try: + link_ade = d[DinkLinkConst.ADE] + return link_ade[DinkLinkConst.ADE_ANIMATEDIFFMODEL] + except KeyError: + if throw_exception: + raise Exception("Could not get AnimateDiffModel class. AnimateDiff-Evolved nodes need to be installed to use SparseCtrl; " + \ + "they are either not installed or are of an insufficient version.") + return None + +def get_AnimateDiffInfo(throw_exception=True): + d = get_dinklink() + try: + link_ade = d[DinkLinkConst.ADE] + return link_ade[DinkLinkConst.ADE_ANIMATEDIFFINFO] + except KeyError: + if throw_exception: + raise Exception("Could not get AnimateDiffInfo class - AnimateDiff-Evolved nodes need to be installed to use SparseCtrl; " + \ + "they are either not installed or are of an insufficient version.") + return None diff --git a/adv_control/sampling.py b/adv_control/sampling.py index d97cd98..7252160 100644 --- a/adv_control/sampling.py +++ b/adv_control/sampling.py @@ -17,9 +17,20 @@ from .control_reference import (ReferenceAdvanced, ReferenceInjections, _forward_inject_BasicTransformerBlock, handle_context_ref_setup, handle_reference_injection, REF_CONTROL_LIST_ALL, CONTEXTREF_CLEAN_FUNC) +from .dinklink import get_dinklink from .utils import torch_dfs, WrapperConsts +def prepare_dinklink_acn_wrapper(): + # expose acn_sampler_sample_wrapper + d = get_dinklink() + link_acn = d.setdefault(WrapperConsts.ACN, {}) + link_acn[WrapperConsts.VERSION] = 10000 + link_acn[WrapperConsts.ACN_CREATE_SAMPLER_SAMPLE_WRAPPER] = (comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, + WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY, + acn_outer_sample_wrapper) + + def support_sliding_context_windows(conds) -> tuple[bool, list[dict]]: # convert to advanced, with report if anything was actually modified modified, new_conds = convert_all_to_advanced(conds)