Refactored DinkLink acn wrapper registration to not cause circular import

This commit is contained in:
Jedrzej Kosinski
2024-11-29 22:40:20 -06:00
parent c59f99efbc
commit e5f7f5a281
4 changed files with 50 additions and 12 deletions
+2
View File
@@ -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()
+1
View File
@@ -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)
+36 -12
View File
@@ -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
+11
View File
@@ -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)