diff --git a/adv_control/control.py b/adv_control/control.py index 4706292..fff8280 100644 --- a/adv_control/control.py +++ b/adv_control/control.py @@ -316,6 +316,7 @@ class SparseCtrlAdvanced(ControlNetAdvanced): self.control_model_wrapped = create_sparse_modelpatcher(self.control_model, load_device=load_device, offload_device=comfy.model_management.unet_offload_device()) self.add_compatible_weight(ControlWeightType.SPARSECTRL) self.control_model: SparseControlNet = self.control_model # does nothing except help with IDE hints + self.postpone_condhint_latents_check = True if self.control_model.use_simplified_conditioning_embedding: # TODO: allow vae_optional to be used instead of preprocessor #self.require_vae = True diff --git a/adv_control/control_lllite.py b/adv_control/control_lllite.py index 76f6d7d..60e0a5f 100644 --- a/adv_control/control_lllite.py +++ b/adv_control/control_lllite.py @@ -302,10 +302,6 @@ class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase): set_model_attn1_patch(model_options, self.patch_attn1.set_control(self)) set_model_attn2_patch(model_options, self.patch_attn2.set_control(self)) - # def patch_model(self, model: ModelPatcher): - # model.set_model_attn1_patch(self.patch_attn1) - # model.set_model_attn2_patch(self.patch_attn2) - def set_cond_hint_inject(self, *args, **kwargs): to_return = super().set_cond_hint_inject(*args, **kwargs) # cond hint for LLLite needs to be scaled between (-1, 1) instead of (0, 1) diff --git a/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index 214b389..616aa70 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -396,6 +396,7 @@ def get_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_name: str raise ValueError(f"No pos_encoder.pe found in SparseCtrl state_dict - {mm_name} is not a valid SparseCtrl model!") +# TODO: replace with DinkLink reference from ADE class SparseCtrlMotionWrapper(nn.Module): def __init__(self, mm_state_dict: dict[str, Tensor], ops=disable_weight_init_clean_groupnorm): super().__init__() diff --git a/adv_control/dinklink.py b/adv_control/dinklink.py index d3874fb..e30a3bd 100644 --- a/adv_control/dinklink.py +++ b/adv_control/dinklink.py @@ -12,6 +12,8 @@ #################################################################################################### from __future__ import annotations import comfy.hooks +from comfy.patcher_extension import WrappersMP + from .sampling import acn_sampler_sample_wrapper DINKLINK = "__DINKLINK" @@ -27,9 +29,12 @@ def get_dinklink() -> dict[str, dict[str]]: class Consts: ACN = "ACN" + VERSION = "version" CREATE_SAMPLER_SAMPLE_WRAPPER = "create_sampler_sample_wrapper" def prepare_dinklink(): # expose acn_sampler_sample_wrapper d = get_dinklink() - d.setdefault(Consts.ACN, {})[Consts.CREATE_SAMPLER_SAMPLE_WRAPPER] = None + link_acn = d.setdefault(Consts.ACN, {}) + link_acn[Consts.VERSION] = 1 + link_acn[Consts.CREATE_SAMPLER_SAMPLE_WRAPPER] = (WrappersMP.SAMPLER_SAMPLE, Consts.ACN, acn_sampler_sample_wrapper) diff --git a/adv_control/nodes.py b/adv_control/nodes.py index 9846f93..a018406 100644 --- a/adv_control/nodes.py +++ b/adv_control/nodes.py @@ -145,11 +145,13 @@ class AdvancedControlNetApply: if is_advanced_controlnet(c_net): # disarm node check c_net.disarm() - # if model required, verify model is passed in, and if so patch it - if c_net.require_model: - if not model_optional: - raise Exception(f"Type '{type(c_net).__name__}' requires model_optional input, but got None.") - c_net.patch_model(model=model_optional) + # check for allow_condhint_latents where vae_optional can't handle it itself + if c_net.allow_condhint_latents and not c_net.require_vae: + if not isinstance(control_hint, AbstractPreprocWrapper): + raise Exception(f"Type '{type(c_net).__name__}' requires proc_IMAGE input via a corresponding preprocessor, but received a normal Image instead.") + else: + if isinstance(control_hint, AbstractPreprocWrapper) and not c_net.postpone_condhint_latents_check: + raise Exception(f"Type '{type(c_net).__name__}' requires a normal Image input, but received a proc_IMAGE input instead.") # if vae required, verify vae is passed in if c_net.require_vae: # if controlnet can accept preprocced condhint latents and is the case, ignore vae requirement diff --git a/adv_control/utils.py b/adv_control/utils.py index d86dfcc..8c98f89 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -457,7 +457,7 @@ class WeightTypeException(TypeError): class AdvancedControlBase: - def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights, require_model=False, require_vae=False, allow_condhint_latents=False): + def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights, require_vae=False, allow_condhint_latents=False): self.base = base self.compatible_weights = [ControlWeightType.UNIVERSAL, ControlWeightType.DEFAULT] self.add_compatible_weight(weights_default.weight_type) @@ -496,14 +496,11 @@ class AdvancedControlBase: # vae to store self.adv_vae = None # require model/vae to be passed into Apply Advanced ControlNet 🛂🅐🅒🅝 node - self.require_model = require_model self.require_vae = require_vae self.allow_condhint_latents = allow_condhint_latents + self.postpone_condhint_latents_check = False # disarm - when set to False, used to force usage of Apply Advanced ControlNet 🛂🅐🅒🅝 node (which will set it to True) - self.disarmed = not require_model - - def patch_model(self, model: ModelPatcher): - pass + self.disarmed = True def add_compatible_weight(self, control_weight_type: str): self.compatible_weights.append(control_weight_type) @@ -874,4 +871,5 @@ class AdvancedControlBase: copied.adv_vae = self.adv_vae copied.require_vae = self.require_vae copied.allow_condhint_latents = self.allow_condhint_latents + copied.postpone_condhint_latents_check = self.postpone_condhint_latents_check copied.disarmed = self.disarmed