Fixed checks for special preprocessor input (ReferenceCN and RGB SparseCtrl), continued work on DinkLink

This commit is contained in:
Jedrzej Kosinski
2024-11-14 03:42:29 -06:00
parent 74320a78e3
commit f42c902fd4
6 changed files with 19 additions and 16 deletions
+1
View File
@@ -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
-4
View File
@@ -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)
+1
View File
@@ -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__()
+6 -1
View File
@@ -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)
+7 -5
View File
@@ -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
+4 -6
View File
@@ -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