Add middle_mult extra via Middle Weight Extras node

This commit is contained in:
Jedrzej Kosinski
2024-11-26 20:21:53 -06:00
parent 8db35a7963
commit dc91f47aed
5 changed files with 45 additions and 11 deletions
+3 -3
View File
@@ -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]]
+2 -2
View File
@@ -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]
+3 -1
View File
@@ -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 🛂🅐🅒🅝",
+28 -1
View File
@@ -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,)
+9 -4
View File
@@ -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):