Fixed checks for special preprocessor input (ReferenceCN and RGB SparseCtrl), continued work on DinkLink
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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__()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user