Add middle_mult extra via Middle Weight Extras node
This commit is contained in:
@@ -14,7 +14,7 @@ from comfy.model_patcher import ModelPatcher
|
||||
from .control_sparsectrl import SparseControlNet, SparseCtrlMotionWrapper, SparseSettings, SparseConst, create_sparse_modelpatcher
|
||||
from .control_lllite import LLLiteModule, LLLitePatch, load_controllllite
|
||||
from .control_svd import svd_unet_config_from_diffusers_unet, SVDControlNet, svd_unet_to_diffusers
|
||||
from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, AbstractPreprocWrapper, ControlWeightType, ControlWeights, WeightTypeException,
|
||||
from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, AbstractPreprocWrapper, ControlWeightType, ControlWeights, WeightTypeException, Extras,
|
||||
manual_cast_clean_groupnorm, disable_weight_init_clean_groupnorm, prepare_mask_batch, get_properly_arranged_t2i_weights, load_torch_file_with_dict_factory,
|
||||
broadcast_image_to_extend, extend_to_batch_size, ORIG_PREVIOUS_CONTROLNET, CONTROL_INIT_BY_ACN)
|
||||
from .logger import logger
|
||||
@@ -30,7 +30,7 @@ class ControlNetAdvanced(ControlNet, AdvancedControlBase):
|
||||
def get_universal_weights(self) -> ControlWeights:
|
||||
def cn_weights_func(idx: int, control: dict[str, list[Tensor]], key: str):
|
||||
if key == "middle":
|
||||
return 1.0
|
||||
return 1.0 * self.weights.extras.get(Extras.MIDDLE_MULT, 1.0)
|
||||
c_len = len(control[key])
|
||||
raw_weights = [(self.weights.base_multiplier ** float((c_len) - i)) for i in range(c_len+1)]
|
||||
raw_weights = raw_weights[:-1]
|
||||
@@ -155,7 +155,7 @@ class T2IAdapterAdvanced(T2IAdapter, AdvancedControlBase):
|
||||
def get_universal_weights(self) -> ControlWeights:
|
||||
def t2i_weights_func(idx: int, control: dict[str, list[Tensor]], key: str):
|
||||
if key == "middle":
|
||||
return 1.0
|
||||
return 1.0 * self.weights.extras.get(Extras.MIDDLE_MULT, 1.0)
|
||||
c_len = 8 #len(control[key])
|
||||
raw_weights = [(self.weights.base_multiplier ** float((c_len-1) - i)) for i in range(c_len)]
|
||||
raw_weights = [raw_weights[-c_len], raw_weights[-3], raw_weights[-2], raw_weights[-1]]
|
||||
|
||||
@@ -22,7 +22,7 @@ import comfy.model_management
|
||||
import comfy.model_detection
|
||||
import comfy.utils
|
||||
|
||||
from .utils import (AdvancedControlBase, ControlWeights, ControlWeightType, TimestepKeyframeGroup, AbstractPreprocWrapper,
|
||||
from .utils import (AdvancedControlBase, ControlWeights, ControlWeightType, TimestepKeyframeGroup, AbstractPreprocWrapper, Extras,
|
||||
extend_to_batch_size, broadcast_image_to_extend)
|
||||
from .logger import logger
|
||||
|
||||
@@ -239,7 +239,7 @@ class ControlNetPlusPlusAdvanced(ControlNet, AdvancedControlBase):
|
||||
def get_universal_weights(self) -> ControlWeights:
|
||||
def cn_weights_func(idx: int, control: dict[str, list[Tensor]], key: str):
|
||||
if key == "middle":
|
||||
return 1.0
|
||||
return 1.0 * self.weights.extras.get(Extras.MIDDLE_MULT, 1.0)
|
||||
c_len = len(control[key])
|
||||
raw_weights = [(self.weights.base_multiplier ** float((c_len) - i)) for i in range(c_len+1)]
|
||||
raw_weights = raw_weights[:-1]
|
||||
|
||||
@@ -4,7 +4,7 @@ from .nodes_main import (ControlNetLoaderAdvanced, DiffControlNetLoaderAdvanced,
|
||||
AdvancedControlNetApply, AdvancedControlNetApplySingle)
|
||||
from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights,
|
||||
SoftControlNetWeightsSD15, CustomControlNetWeightsSD15, CustomControlNetWeightsFlux,
|
||||
SoftT2IAdapterWeights, CustomT2IAdapterWeights)
|
||||
SoftT2IAdapterWeights, CustomT2IAdapterWeights, ExtrasMiddleMultNode)
|
||||
from .nodes_keyframes import (LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode,
|
||||
TimestepKeyframeNode, TimestepKeyframeInterpolationNode, TimestepKeyframeFromStrengthListNode)
|
||||
from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor, SparseWeightExtras
|
||||
@@ -45,6 +45,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ACN_SoftT2IAdapterWeights": SoftT2IAdapterWeights,
|
||||
"ACN_CustomT2IAdapterWeights": CustomT2IAdapterWeights,
|
||||
"ACN_DefaultUniversalWeights": DefaultWeights,
|
||||
"ACN_ExtrasMiddleMult": ExtrasMiddleMultNode,
|
||||
# SparseCtrl
|
||||
"ACN_SparseCtrlRGBPreprocessor": RgbSparseCtrlPreprocessor,
|
||||
"ACN_SparseCtrlLoaderAdvanced": SparseCtrlLoaderAdvanced,
|
||||
@@ -101,6 +102,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ACN_SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝",
|
||||
"ACN_CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝",
|
||||
"ACN_DefaultUniversalWeights": "Default Weights 🛂🅐🅒🅝",
|
||||
"ACN_ExtrasMiddleMult": "Middle Weight Extras 🛂🅐🅒🅝",
|
||||
# SparseCtrl
|
||||
"ACN_SparseCtrlRGBPreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝",
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from torch import Tensor
|
||||
import torch
|
||||
from .utils import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion
|
||||
from .utils import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, Extras, get_properly_arranged_t2i_weights, linear_conversion
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -300,3 +300,30 @@ class CustomT2IAdapterWeights:
|
||||
weights = get_properly_arranged_t2i_weights(weights)
|
||||
weights = ControlWeights.t2iadapter(weights_input=weights, uncond_multiplier=uncond_multiplier, extras=cn_extras, disable_applied_to=True)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class ExtrasMiddleMultNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"middle_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
|
||||
},
|
||||
"optional": {
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CN_WEIGHTS_EXTRAS",)
|
||||
RETURN_NAMES = ("cn_extras",)
|
||||
FUNCTION = "create_extras"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/extras"
|
||||
|
||||
def create_extras(self, middle_mult: float, cn_extras: dict[str]={}):
|
||||
cn_extras = cn_extras.copy()
|
||||
cn_extras[Extras.MIDDLE_MULT] = middle_mult
|
||||
return (cn_extras,)
|
||||
|
||||
@@ -22,6 +22,9 @@ BIGMAX = (2**53-1)
|
||||
ORIG_PREVIOUS_CONTROLNET = "_orig_previous_controlnet"
|
||||
CONTROL_INIT_BY_ACN = "_control_init_by_ACN"
|
||||
|
||||
class Extras:
|
||||
MIDDLE_MULT = "middle_mult"
|
||||
|
||||
|
||||
def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_load_torch_file: Callable):
|
||||
def load_torch_file_with_dict(*args, **kwargs):
|
||||
@@ -77,17 +80,19 @@ class ControlWeights:
|
||||
self.has_uncond_multiplier = not math.isclose(self.uncond_multiplier, 1.0)
|
||||
self.uncond_mask = uncond_mask if uncond_mask is not None else 1.0
|
||||
self.has_uncond_mask = uncond_mask is not None
|
||||
self.extras = extras
|
||||
self.extras = extras.copy()
|
||||
self.disable_applied_to = disable_applied_to
|
||||
|
||||
def get(self, idx: int, control: dict[str, list[Tensor]], key: str, default=1.0) -> Union[float, Tensor]:
|
||||
# if weight_func present, use it
|
||||
if self.weight_func is not None:
|
||||
return self.weight_func(idx=idx, control=control, key=key)
|
||||
effective_mult = 1.0
|
||||
# if weights is not none, return index
|
||||
relevant_weights = None
|
||||
if key == "middle":
|
||||
relevant_weights = self.weights_middle
|
||||
effective_mult *= self.extras.get(Extras.MIDDLE_MULT, 1.0)
|
||||
elif key == "input":
|
||||
relevant_weights = self.weights_input
|
||||
if relevant_weights is not None:
|
||||
@@ -95,10 +100,10 @@ class ControlWeights:
|
||||
else:
|
||||
relevant_weights = self.weights_output
|
||||
if relevant_weights is None:
|
||||
return default
|
||||
return default * effective_mult
|
||||
elif idx >= len(relevant_weights):
|
||||
return default
|
||||
return relevant_weights[idx]
|
||||
return default * effective_mult
|
||||
return relevant_weights[idx] * effective_mult
|
||||
|
||||
def copy_with_new_weights(self, new_weights_input: list[float]=None, new_weights_middle: list[float]=None, new_weights_output: list[float]=None,
|
||||
new_weight_func: Callable=None):
|
||||
|
||||
Reference in New Issue
Block a user