Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ed5deb30de |
@@ -7,18 +7,14 @@ on:
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'Kosinkadink' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
with:
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }} ## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }} ## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
@@ -12,10 +12,7 @@ ControlNet preprocessors are available through [comfyui_controlnet_aux](https://
|
||||
- Replicate ***"ControlNet is more important"*** feature from sd-webui-controlnet extension via ***uncond_multiplier*** on ***Soft Weights***
|
||||
- uncond_multiplier=0.0 gives identical results of auto1111's feature, but values between 0.0 and 1.0 can be used without issue to granularly control the setting.
|
||||
- ControlNet, T2IAdapter, and ControlLoRA support for sliding context windows
|
||||
- ControlLLLite support
|
||||
- ControlNet++ support
|
||||
- CtrLoRA support
|
||||
- Relevant models linked on [CtrLoRA github page](https://github.com/xyfJASON/ctrlora)
|
||||
- ControlLLLite support (requires model_optional to be passed into and out of Apply Advanced ControlNet node)
|
||||
- SparseCtrl support
|
||||
- SVD-ControlNet support
|
||||
- Stable Video Diffusion ControlNets trained by **CiaraRowles**: [Depth](https://huggingface.co/CiaraRowles/temporal-controlnet-depth-svd-v1/tree/main/controlnet), [Lineart](https://huggingface.co/CiaraRowles/temporal-controlnet-lineart-svd-v1/tree/main/controlnet)
|
||||
|
||||
+1
-9
@@ -1,11 +1,3 @@
|
||||
from .adv_control.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .adv_control import documentation
|
||||
from .adv_control.dinklink import init_dinklink
|
||||
from .adv_control.sampling import prepare_dinklink_acn_wrapper
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', "WEB_DIRECTORY"]
|
||||
documentation.format_descriptions(NODE_CLASS_MAPPINGS)
|
||||
|
||||
init_dinklink()
|
||||
prepare_dinklink_acn_wrapper()
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
+392
-509
File diff suppressed because it is too large
Load Diff
@@ -1,231 +0,0 @@
|
||||
# Core code adapted from CtrLoRA github repo:
|
||||
# https://github.com/xyfJASON/ctrlora
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from comfy.cldm.cldm import ControlNet as ControlNetCLDM
|
||||
import comfy.model_detection
|
||||
import comfy.model_management
|
||||
import comfy.ops
|
||||
import comfy.utils
|
||||
|
||||
from comfy.ldm.modules.diffusionmodules.util import (
|
||||
zero_module,
|
||||
timestep_embedding,
|
||||
)
|
||||
|
||||
from .control import ControlNetAdvanced
|
||||
from .utils import TimestepKeyframeGroup
|
||||
from .logger import logger
|
||||
|
||||
|
||||
class ControlNetCtrLoRA(ControlNetCLDM):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
# delete input hint block
|
||||
del self.input_hint_block
|
||||
|
||||
def forward(self, x: Tensor, hint: Tensor, timesteps, context, y=None, **kwargs):
|
||||
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype)
|
||||
emb = self.time_embed(t_emb)
|
||||
|
||||
out_output = []
|
||||
out_middle = []
|
||||
|
||||
if self.num_classes is not None:
|
||||
assert y.shape[0] == x.shape[0]
|
||||
emb = emb + self.label_emb(y)
|
||||
|
||||
h = hint.to(dtype=x.dtype)
|
||||
for module, zero_conv in zip(self.input_blocks, self.zero_convs):
|
||||
h = module(h, emb, context)
|
||||
out_output.append(zero_conv(h, emb, context))
|
||||
|
||||
h = self.middle_block(h, emb, context)
|
||||
out_middle.append(self.middle_block_out(h, emb, context))
|
||||
|
||||
return {"middle": out_middle, "output": out_output}
|
||||
|
||||
|
||||
class CtrLoRAAdvanced(ControlNetAdvanced):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.preprocess_image = lambda a: (a + 1) / 2.0
|
||||
self.require_vae = True
|
||||
self.mult_by_ratio_when_vae = False
|
||||
|
||||
def pre_run_advanced(self, model, percent_to_timestep_function):
|
||||
super().pre_run_advanced(model, percent_to_timestep_function)
|
||||
self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond hint
|
||||
|
||||
def cleanup_advanced(self):
|
||||
super().cleanup_advanced()
|
||||
if self.latent_format is not None:
|
||||
del self.latent_format
|
||||
self.latent_format = None
|
||||
|
||||
def copy(self):
|
||||
c = CtrLoRAAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling, load_device=self.load_device, manual_cast_dtype=self.manual_cast_dtype)
|
||||
c.control_model = self.control_model
|
||||
c.control_model_wrapped = self.control_model_wrapped
|
||||
self.copy_to(c)
|
||||
self.copy_to_advanced(c)
|
||||
return c
|
||||
|
||||
|
||||
def load_ctrlora(base_path: str, lora_path: str,
|
||||
base_data: dict[str, Tensor]=None, lora_data: dict[str, Tensor]=None,
|
||||
timestep_keyframe: TimestepKeyframeGroup=None, model=None, model_options={}):
|
||||
if base_data is None:
|
||||
base_data = comfy.utils.load_torch_file(base_path, safe_load=True)
|
||||
controlnet_data = base_data
|
||||
|
||||
# first, check that base_data contains keys with lora_layer
|
||||
contains_lora_layers = False
|
||||
for key in base_data:
|
||||
if "lora_layer" in key:
|
||||
contains_lora_layers = True
|
||||
if not contains_lora_layers:
|
||||
raise Exception(f"File '{base_path}' is not a valid CtrLoRA base model; does not contain any lora_layer keys.")
|
||||
|
||||
controlnet_config = None
|
||||
supported_inference_dtypes = None
|
||||
|
||||
pth_key = 'control_model.zero_convs.0.0.weight'
|
||||
pth = False
|
||||
key = 'zero_convs.0.0.weight'
|
||||
if pth_key in controlnet_data:
|
||||
pth = True
|
||||
key = pth_key
|
||||
prefix = "control_model."
|
||||
elif key in controlnet_data:
|
||||
prefix = ""
|
||||
else:
|
||||
raise Exception("")
|
||||
net = load_t2i_adapter(controlnet_data, model_options=model_options)
|
||||
if net is None:
|
||||
logging.error("error could not detect control model type.")
|
||||
return net
|
||||
|
||||
if controlnet_config is None:
|
||||
model_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, True)
|
||||
supported_inference_dtypes = list(model_config.supported_inference_dtypes)
|
||||
controlnet_config = model_config.unet_config
|
||||
|
||||
unet_dtype = model_options.get("dtype", None)
|
||||
if unet_dtype is None:
|
||||
weight_dtype = comfy.utils.weight_dtype(controlnet_data)
|
||||
|
||||
if supported_inference_dtypes is None:
|
||||
supported_inference_dtypes = [comfy.model_management.unet_dtype()]
|
||||
|
||||
if weight_dtype is not None:
|
||||
supported_inference_dtypes.append(weight_dtype)
|
||||
|
||||
unet_dtype = comfy.model_management.unet_dtype(model_params=-1, supported_dtypes=supported_inference_dtypes)
|
||||
|
||||
load_device = comfy.model_management.get_torch_device()
|
||||
|
||||
manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device)
|
||||
operations = model_options.get("custom_operations", None)
|
||||
if operations is None:
|
||||
operations = comfy.ops.pick_operations(unet_dtype, manual_cast_dtype)
|
||||
|
||||
controlnet_config["operations"] = operations
|
||||
controlnet_config["dtype"] = unet_dtype
|
||||
controlnet_config["device"] = comfy.model_management.unet_offload_device()
|
||||
controlnet_config.pop("out_channels")
|
||||
controlnet_config["hint_channels"] = 3
|
||||
#controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1]
|
||||
control_model = ControlNetCtrLoRA(**controlnet_config)
|
||||
|
||||
if pth:
|
||||
if 'difference' in controlnet_data:
|
||||
if model is not None:
|
||||
comfy.model_management.load_models_gpu([model])
|
||||
model_sd = model.model_state_dict()
|
||||
for x in controlnet_data:
|
||||
c_m = "control_model."
|
||||
if x.startswith(c_m):
|
||||
sd_key = "diffusion_model.{}".format(x[len(c_m):])
|
||||
if sd_key in model_sd:
|
||||
cd = controlnet_data[x]
|
||||
cd += model_sd[sd_key].type(cd.dtype).to(cd.device)
|
||||
else:
|
||||
logger.warning("WARNING: Loaded a diff controlnet without a model. It will very likely not work.")
|
||||
|
||||
class WeightsLoader(torch.nn.Module):
|
||||
pass
|
||||
w = WeightsLoader()
|
||||
w.control_model = control_model
|
||||
missing, unexpected = w.load_state_dict(controlnet_data, strict=False)
|
||||
else:
|
||||
missing, unexpected = control_model.load_state_dict(controlnet_data, strict=False)
|
||||
|
||||
if len(missing) > 0:
|
||||
logger.warning("missing controlnet keys: {}".format(missing))
|
||||
|
||||
if len(unexpected) > 0:
|
||||
logger.debug("unexpected controlnet keys: {}".format(unexpected))
|
||||
|
||||
global_average_pooling = model_options.get("global_average_pooling", False)
|
||||
control = CtrLoRAAdvanced(control_model, timestep_keyframe, global_average_pooling=global_average_pooling,
|
||||
load_device=load_device, manual_cast_dtype=manual_cast_dtype)
|
||||
# load lora data onto the controlnet
|
||||
if lora_path is not None:
|
||||
load_lora_data(control, lora_path)
|
||||
|
||||
return control
|
||||
|
||||
|
||||
def load_lora_data(control: CtrLoRAAdvanced, lora_path: str, loaded_data: dict[str, Tensor]=None, lora_strength=1.0):
|
||||
if loaded_data is None:
|
||||
loaded_data = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||
# check that lora_data contains keys with lora_layer
|
||||
contains_lora_layers = False
|
||||
for key in loaded_data:
|
||||
if "lora_layer" in key:
|
||||
contains_lora_layers = True
|
||||
if not contains_lora_layers:
|
||||
raise Exception(f"File '{lora_path}' is not a valid CtrLoRA lora model; does not contain any lora_layer keys.")
|
||||
|
||||
# now that we know we have a ctrlora file, separate keys into 'set' and 'lora' keys
|
||||
data_set: dict[str, Tensor] = {}
|
||||
data_lora: dict[str, Tensor] = {}
|
||||
|
||||
for key in list(loaded_data.keys()):
|
||||
if 'lora_layer' in key:
|
||||
data_lora[key] = loaded_data.pop(key)
|
||||
else:
|
||||
data_set[key] = loaded_data.pop(key)
|
||||
# no keys should be left over
|
||||
if len(loaded_data) > 0:
|
||||
logger.warning("Not all keys from CtrlLoRA lora model's loaded data were parsed!")
|
||||
|
||||
# turn set/lora data into corresponding patches;
|
||||
patches = {}
|
||||
# set will replace the values
|
||||
for key, value in data_set.items():
|
||||
# prase model key from key;
|
||||
# remove "control_model."
|
||||
model_key = key.replace("control_model.", "")
|
||||
patches[model_key] = ("set", (value,))
|
||||
# lora will do mm of up and down tensors
|
||||
for down_key in data_lora:
|
||||
# only process lora down keys; we will process both up+down at the same time
|
||||
if ".up." in down_key:
|
||||
continue
|
||||
# get up version of down key
|
||||
up_key = down_key.replace(".down.", ".up.")
|
||||
# get key that will match up with model key;
|
||||
# remove "lora_layer.down." and "control_model."
|
||||
model_key = down_key.replace("lora_layer.down.", "").replace("control_model.", "")
|
||||
|
||||
weight_down = data_lora[down_key]
|
||||
weight_up = data_lora[up_key]
|
||||
# currently, ComfyUI expects 6 elements in 'lora' type, but for future-proofing add a bunch more with None
|
||||
patches[model_key] = ("lora", (weight_up, weight_down, None, None, None, None,
|
||||
None, None, None, None, None, None, None, None))
|
||||
|
||||
# now that patches are made, add them to model
|
||||
control.control_model_wrapped.add_patches(patches, strength_patch=lora_strength)
|
||||
+19
-192
@@ -8,33 +8,10 @@ import torch
|
||||
import os
|
||||
|
||||
import comfy.utils
|
||||
import comfy.ops
|
||||
import comfy.model_management
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.controlnet import ControlBase
|
||||
|
||||
from .logger import logger
|
||||
from .utils import (AdvancedControlBase, TimestepKeyframeGroup, ControlWeights, broadcast_image_to_extend, extend_to_batch_size,
|
||||
prepare_mask_batch)
|
||||
|
||||
|
||||
# based on set_model_patch code in comfy/model_patcher.py
|
||||
def set_model_patch(transformer_options, patch, name):
|
||||
to = transformer_options
|
||||
# check if patch was already added
|
||||
if "patches" in to:
|
||||
current_patches = to["patches"].get(name, [])
|
||||
if patch in current_patches:
|
||||
return
|
||||
if "patches" not in to:
|
||||
to["patches"] = {}
|
||||
to["patches"][name] = to["patches"].get(name, []) + [patch]
|
||||
|
||||
def set_model_attn1_patch(transformer_options, patch):
|
||||
set_model_patch(transformer_options, patch, "attn1_patch")
|
||||
|
||||
def set_model_attn2_patch(transformer_options, patch):
|
||||
set_model_patch(transformer_options, patch, "attn2_patch")
|
||||
from .utils import AdvancedControlBase, deepcopy_with_sharing, prepare_mask_batch
|
||||
|
||||
|
||||
def extra_options_to_module_prefix(extra_options):
|
||||
@@ -115,8 +92,26 @@ class LLLitePatch:
|
||||
return LLLitePatch(self.modules, self.patch_type, control)
|
||||
|
||||
def cleanup(self):
|
||||
#total_cleaned = 0
|
||||
for module in self.modules.values():
|
||||
module.cleanup()
|
||||
# total_cleaned += 1
|
||||
#logger.info(f"cleaned modules: {total_cleaned}, {id(self)}")
|
||||
#logger.error(f"cleanup LLLitePatch: {id(self)}")
|
||||
|
||||
# make sure deepcopy does not copy control, and deepcopied LLLitePatch should be assigned to control
|
||||
def __deepcopy__(self, memo):
|
||||
self.cleanup()
|
||||
to_return: LLLitePatch = deepcopy_with_sharing(self, shared_attribute_names = ['control'], memo=memo)
|
||||
#logger.warn(f"patch {id(self)} turned into {id(to_return)}")
|
||||
try:
|
||||
if self.patch_type == self.ATTN1:
|
||||
to_return.control.patch_attn1 = to_return
|
||||
elif self.patch_type == self.ATTN2:
|
||||
to_return.control.patch_attn2 = to_return
|
||||
except Exception:
|
||||
pass
|
||||
return to_return
|
||||
|
||||
|
||||
# TODO: use comfy.ops to support fp8 properly
|
||||
@@ -257,171 +252,3 @@ class LLLiteModule(torch.nn.Module):
|
||||
if cond_type == 1:
|
||||
cx[actual_length*idx:actual_length*(idx+1)] *= control.weights.uncond_multiplier
|
||||
return cx * mask * control.strength * control._current_timestep_keyframe.strength
|
||||
|
||||
|
||||
class ControlLLLiteModules(torch.nn.Module):
|
||||
def __init__(self, patch_attn1: LLLitePatch, patch_attn2: LLLitePatch):
|
||||
super().__init__()
|
||||
self.patch_attn1_modules = torch.nn.Sequential(*list(patch_attn1.modules.values()))
|
||||
self.patch_attn2_modules = torch.nn.Sequential(*list(patch_attn2.modules.values()))
|
||||
|
||||
|
||||
class ControlLLLiteAdvanced(ControlBase, AdvancedControlBase):
|
||||
# This ControlNet is more of an attention patch than a traditional controlnet
|
||||
def __init__(self, patch_attn1: LLLitePatch, patch_attn2: LLLitePatch, timestep_keyframes: TimestepKeyframeGroup, device, ops: comfy.ops.disable_weight_init):
|
||||
super().__init__()
|
||||
AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite())
|
||||
self.device = device
|
||||
self.ops = ops
|
||||
self.patch_attn1 = patch_attn1.clone_with_control(self)
|
||||
self.patch_attn2 = patch_attn2.clone_with_control(self)
|
||||
self.control_model = ControlLLLiteModules(self.patch_attn1, self.patch_attn2)
|
||||
self.control_model_wrapped = ModelPatcher(self.control_model, load_device=device, offload_device=comfy.model_management.unet_offload_device())
|
||||
self.latent_dims_div2 = None
|
||||
self.latent_dims_div4 = None
|
||||
|
||||
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)
|
||||
self.cond_hint_original = self.cond_hint_original * 2.0 - 1.0
|
||||
return to_return
|
||||
|
||||
def pre_run_advanced(self, *args, **kwargs):
|
||||
AdvancedControlBase.pre_run_advanced(self, *args, **kwargs)
|
||||
#logger.error(f"in cn: {id(self.patch_attn1)},{id(self.patch_attn2)}")
|
||||
self.patch_attn1.set_control(self)
|
||||
self.patch_attn2.set_control(self)
|
||||
#logger.warn(f"in pre_run_advanced: {id(self)}")
|
||||
|
||||
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int, transformer_options: dict):
|
||||
# normal ControlNet stuff
|
||||
control_prev = None
|
||||
if self.previous_controlnet is not None:
|
||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number, transformer_options)
|
||||
|
||||
if self.timestep_range is not None:
|
||||
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
|
||||
return control_prev
|
||||
|
||||
dtype = x_noisy.dtype
|
||||
# prepare cond_hint
|
||||
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * 8 != self.cond_hint.shape[2] or x_noisy.shape[3] * 8 != self.cond_hint.shape[3]:
|
||||
if self.cond_hint is not None:
|
||||
del self.cond_hint
|
||||
self.cond_hint = None
|
||||
# if self.cond_hint_original length greater or equal to real latent count, subdivide it before scaling
|
||||
if self.sub_idxs is not None:
|
||||
actual_cond_hint_orig = self.cond_hint_original
|
||||
if self.cond_hint_original.size(0) < self.full_latent_length:
|
||||
actual_cond_hint_orig = extend_to_batch_size(tensor=actual_cond_hint_orig, batch_size=self.full_latent_length)
|
||||
self.cond_hint = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(x_noisy.device)
|
||||
else:
|
||||
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(x_noisy.device)
|
||||
if x_noisy.shape[0] != self.cond_hint.shape[0]:
|
||||
self.cond_hint = broadcast_image_to_extend(self.cond_hint, x_noisy.shape[0], batched_number)
|
||||
# some special logic here compared to other controlnets:
|
||||
# * The cond_emb in attn patches will divide latent dims by 2 or 4, integer
|
||||
# * Due to this loss, the cond_emb will become smaller than x input if latent dims are not divisble by 2 or 4
|
||||
divisible_by_2_h = x_noisy.shape[2]%2==0
|
||||
divisible_by_2_w = x_noisy.shape[3]%2==0
|
||||
if not (divisible_by_2_h and divisible_by_2_w):
|
||||
#logger.warn(f"{x_noisy.shape} not divisible by 2!")
|
||||
new_h = (x_noisy.shape[2]//2)*2
|
||||
new_w = (x_noisy.shape[3]//2)*2
|
||||
if not divisible_by_2_h:
|
||||
new_h += 2
|
||||
if not divisible_by_2_w:
|
||||
new_w += 2
|
||||
self.latent_dims_div2 = (new_h, new_w)
|
||||
divisible_by_4_h = x_noisy.shape[2]%4==0
|
||||
divisible_by_4_w = x_noisy.shape[3]%4==0
|
||||
if not (divisible_by_4_h and divisible_by_4_w):
|
||||
#logger.warn(f"{x_noisy.shape} not divisible by 4!")
|
||||
new_h = (x_noisy.shape[2]//4)*4
|
||||
new_w = (x_noisy.shape[3]//4)*4
|
||||
if not divisible_by_4_h:
|
||||
new_h += 4
|
||||
if not divisible_by_4_w:
|
||||
new_w += 4
|
||||
self.latent_dims_div4 = (new_h, new_w)
|
||||
# prepare mask
|
||||
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number)
|
||||
# done preparing; model patches will take care of everything now
|
||||
set_model_attn1_patch(transformer_options, self.patch_attn1.set_control(self))
|
||||
set_model_attn2_patch(transformer_options, self.patch_attn2.set_control(self))
|
||||
# return normal controlnet stuff
|
||||
return control_prev
|
||||
|
||||
def get_models(self):
|
||||
to_return: list = super().get_models()
|
||||
to_return.append(self.control_model_wrapped)
|
||||
return to_return
|
||||
|
||||
def cleanup_advanced(self):
|
||||
super().cleanup_advanced()
|
||||
self.patch_attn1.cleanup()
|
||||
self.patch_attn2.cleanup()
|
||||
self.latent_dims_div2 = None
|
||||
self.latent_dims_div4 = None
|
||||
|
||||
def copy(self):
|
||||
c = ControlLLLiteAdvanced(self.patch_attn1, self.patch_attn2, self.timestep_keyframes, self.device, self.ops)
|
||||
self.copy_to(c)
|
||||
self.copy_to_advanced(c)
|
||||
return c
|
||||
|
||||
|
||||
def load_controllllite(ckpt_path: str, controlnet_data: dict[str, Tensor]=None, timestep_keyframe: TimestepKeyframeGroup=None):
|
||||
if controlnet_data is None:
|
||||
controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
|
||||
# adapted from https://github.com/kohya-ss/ControlNet-LLLite-ComfyUI
|
||||
# first, split weights for each module
|
||||
module_weights = {}
|
||||
for key, value in controlnet_data.items():
|
||||
fragments = key.split(".")
|
||||
module_name = fragments[0]
|
||||
weight_name = ".".join(fragments[1:])
|
||||
|
||||
if module_name not in module_weights:
|
||||
module_weights[module_name] = {}
|
||||
module_weights[module_name][weight_name] = value
|
||||
|
||||
unet_dtype = comfy.model_management.unet_dtype()
|
||||
load_device = comfy.model_management.get_torch_device()
|
||||
manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device)
|
||||
ops = comfy.ops.disable_weight_init
|
||||
if manual_cast_dtype is not None:
|
||||
ops = comfy.ops.manual_cast
|
||||
|
||||
# next, load each module
|
||||
modules = {}
|
||||
for module_name, weights in module_weights.items():
|
||||
# kohya planned to do something about how these should be chosen, so I'm not touching this
|
||||
# since I am not familiar with the logic for this
|
||||
if "conditioning1.4.weight" in weights:
|
||||
depth = 3
|
||||
elif weights["conditioning1.2.weight"].shape[-1] == 4:
|
||||
depth = 2
|
||||
else:
|
||||
depth = 1
|
||||
|
||||
module = LLLiteModule(
|
||||
name=module_name,
|
||||
is_conv2d=weights["down.0.weight"].ndim == 4,
|
||||
in_dim=weights["down.0.weight"].shape[1],
|
||||
depth=depth,
|
||||
cond_emb_dim=weights["conditioning1.0.weight"].shape[0] * 2,
|
||||
mlp_dim=weights["down.0.weight"].shape[0],
|
||||
)
|
||||
# load weights into module
|
||||
module.load_state_dict(weights)
|
||||
modules[module_name] = module.to(dtype=unet_dtype)
|
||||
if len(modules) == 1:
|
||||
module.is_first = True
|
||||
|
||||
#logger.info(f"loaded {ckpt_path} successfully, {len(modules)} modules")
|
||||
|
||||
patch_attn1 = LLLitePatch(modules=modules, patch_type=LLLitePatch.ATTN1)
|
||||
patch_attn2 = LLLitePatch(modules=modules, patch_type=LLLitePatch.ATTN2)
|
||||
control = ControlLLLiteAdvanced(patch_attn1=patch_attn1, patch_attn2=patch_attn2, timestep_keyframes=timestep_keyframe, device=load_device, ops=ops)
|
||||
return control
|
||||
|
||||
@@ -1,486 +0,0 @@
|
||||
# Code ported and modified from the diffusers ControlNetPlus repo by Qi Xin:
|
||||
# https://github.com/xinsir6/ControlNetPlus/blob/main/models/controlnet_union.py
|
||||
from typing import Union
|
||||
|
||||
import os
|
||||
import torch
|
||||
import torch as th
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
from collections import OrderedDict
|
||||
|
||||
|
||||
from comfy.ldm.modules.diffusionmodules.util import (zero_module, timestep_embedding)
|
||||
|
||||
from comfy.cldm.cldm import ControlNet as ControlNetCLDM
|
||||
import comfy.cldm.cldm
|
||||
from comfy.controlnet import ControlNet
|
||||
#from comfy.t2i_adapter.adapter import ResidualAttentionBlock
|
||||
from comfy.ldm.modules.attention import optimized_attention
|
||||
import comfy.ops
|
||||
import comfy.model_base
|
||||
import comfy.model_management
|
||||
import comfy.model_detection
|
||||
import comfy.utils
|
||||
|
||||
from .utils import (AdvancedControlBase, ControlWeights, ControlWeightType, TimestepKeyframeGroup, AbstractPreprocWrapper, Extras,
|
||||
extend_to_batch_size, broadcast_image_to_extend)
|
||||
from .logger import logger
|
||||
|
||||
|
||||
class PlusPlusType:
|
||||
OPENPOSE = "openpose"
|
||||
DEPTH = "depth"
|
||||
THICKLINE = "hed/pidi/scribble/ted"
|
||||
THINLINE = "canny/lineart/mlsd"
|
||||
NORMAL = "normal"
|
||||
SEGMENT = "segment"
|
||||
TILE = "tile"
|
||||
REPAINT = "inpaint/outpaint"
|
||||
NONE = "none"
|
||||
_LIST_WITH_NONE = [OPENPOSE, DEPTH, THICKLINE, THINLINE, NORMAL, SEGMENT, TILE, REPAINT, NONE]
|
||||
_LIST = [OPENPOSE, DEPTH, THICKLINE, THINLINE, NORMAL, SEGMENT, TILE, REPAINT]
|
||||
_DICT = {OPENPOSE: 0, DEPTH: 1, THICKLINE: 2, THINLINE: 3, NORMAL: 4, SEGMENT: 5, TILE: 6, REPAINT: 7, NONE: -1}
|
||||
|
||||
@classmethod
|
||||
def to_idx(cls, control_type: str):
|
||||
try:
|
||||
return cls._DICT[control_type]
|
||||
except KeyError:
|
||||
raise Exception(f"Unknown control type '{control_type}'.")
|
||||
|
||||
|
||||
class PlusPlusInput:
|
||||
def __init__(self, image: Tensor, control_type: str, strength: float):
|
||||
self.image = image
|
||||
self.control_type = control_type
|
||||
self.strength = strength
|
||||
|
||||
def clone(self):
|
||||
return PlusPlusInput(self.image, self.control_type, self.strength)
|
||||
|
||||
|
||||
class PlusPlusInputGroup:
|
||||
def __init__(self):
|
||||
self.controls: dict[str, PlusPlusInput] = {}
|
||||
|
||||
def add(self, pp_input: PlusPlusInput):
|
||||
if pp_input.control_type in self.controls:
|
||||
raise Exception(f"Control type '{pp_input.control_type}' is already present; ControlNet++ does not allow more than 1 of each type.")
|
||||
self.controls[pp_input.control_type] = pp_input
|
||||
|
||||
def clone(self) -> 'PlusPlusInputGroup':
|
||||
cloned = PlusPlusInputGroup()
|
||||
for key, value in self.controls.items():
|
||||
cloned.controls[key] = value.clone()
|
||||
return cloned
|
||||
|
||||
|
||||
class PlusPlusImageWrapper(AbstractPreprocWrapper):
|
||||
error_msg = error_msg = "Invalid use of ControlNet++ Image Wrapper. The output of ControlNet++ Image Wrapper is NOT a usual image, but an object holding the images and extra info - you must connect the output directly to an Apply Advanced ControlNet node. It cannot be used for anything else that accepts IMAGE input."
|
||||
def __init__(self, condhint: PlusPlusInputGroup):
|
||||
super().__init__(condhint)
|
||||
# just an IDE type hint
|
||||
self.condhint: PlusPlusInputGroup
|
||||
|
||||
def movedim(self, source: int, destination: int):
|
||||
condhint = self.condhint.clone()
|
||||
for pp_input in condhint.controls.values():
|
||||
pp_input.image = pp_input.image.movedim(source, destination)
|
||||
return PlusPlusImageWrapper(condhint)
|
||||
|
||||
# parts taken from comfy/cldm/cldm.py
|
||||
class OptimizedAttention(nn.Module):
|
||||
def __init__(self, c, nhead, dropout=0.0, dtype=None, device=None, operations=None):
|
||||
super().__init__()
|
||||
self.heads = nhead
|
||||
self.c = c
|
||||
|
||||
self.in_proj = operations.Linear(c, c * 3, bias=True, dtype=dtype, device=device)
|
||||
self.out_proj = operations.Linear(c, c, bias=True, dtype=dtype, device=device)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.in_proj(x)
|
||||
q, k, v = x.split(self.c, dim=2)
|
||||
out = optimized_attention(q, k, v, self.heads)
|
||||
return self.out_proj(out)
|
||||
|
||||
class QuickGELU(nn.Module):
|
||||
def forward(self, x: torch.Tensor):
|
||||
return x * torch.sigmoid(1.702 * x)
|
||||
|
||||
class ResBlockUnionControlnet(nn.Module):
|
||||
def __init__(self, dim, nhead, dtype=None, device=None, operations=None):
|
||||
super().__init__()
|
||||
self.attn = OptimizedAttention(dim, nhead, dtype=dtype, device=device, operations=operations)
|
||||
self.ln_1 = operations.LayerNorm(dim, dtype=dtype, device=device)
|
||||
self.mlp = nn.Sequential(
|
||||
OrderedDict([("c_fc", operations.Linear(dim, dim * 4, dtype=dtype, device=device)), ("gelu", QuickGELU()),
|
||||
("c_proj", operations.Linear(dim * 4, dim, dtype=dtype, device=device))]))
|
||||
self.ln_2 = operations.LayerNorm(dim, dtype=dtype, device=device)
|
||||
|
||||
def attention(self, x: torch.Tensor):
|
||||
return self.attn(x)
|
||||
|
||||
def forward(self, x: torch.Tensor):
|
||||
x = x + self.attention(self.ln_1(x))
|
||||
x = x + self.mlp(self.ln_2(x))
|
||||
return x
|
||||
|
||||
|
||||
class ControlAddEmbeddingAdv(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, num_control_type, dtype=None, device=None, operations: comfy.ops.disable_weight_init=None):
|
||||
super().__init__()
|
||||
self.num_control_type = num_control_type
|
||||
self.in_dim = in_dim
|
||||
self.linear_1 = operations.Linear(in_dim * num_control_type, out_dim, dtype=dtype, device=device)
|
||||
self.linear_2 = operations.Linear(out_dim, out_dim, dtype=dtype, device=device)
|
||||
|
||||
def forward(self, control_type, dtype, device):
|
||||
if control_type is None:
|
||||
control_type = torch.zeros((self.num_control_type,), device=device)
|
||||
c_type = timestep_embedding(control_type.flatten(), self.in_dim, repeat_only=False).to(dtype).reshape((-1, self.num_control_type * self.in_dim))
|
||||
return self.linear_2(torch.nn.functional.silu(self.linear_1(c_type)))
|
||||
|
||||
|
||||
class ControlNetPlusPlus(ControlNetCLDM):
|
||||
def __init__(self, *args,**kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
operations: comfy.ops.disable_weight_init = kwargs.get("operations", comfy.ops.disable_weight_init)
|
||||
device = kwargs.get("device", None)
|
||||
|
||||
time_embed_dim = self.model_channels * 4
|
||||
control_add_embed_dim = 256
|
||||
|
||||
self.control_add_embedding = ControlAddEmbeddingAdv(control_add_embed_dim, time_embed_dim, self.num_control_type, dtype=self.dtype, device=device, operations=operations)
|
||||
|
||||
def union_controlnet_merge(self, hint: list[Tensor], control_type, emb, context):
|
||||
# Equivalent to: https://github.com/xinsir6/ControlNetPlus/tree/main
|
||||
indexes = torch.nonzero(control_type[0])
|
||||
inputs = []
|
||||
condition_list = []
|
||||
|
||||
for idx in range(indexes.shape[0]):
|
||||
controlnet_cond = self.input_hint_block(hint[indexes[idx][0]], emb, context)
|
||||
feat_seq = torch.mean(controlnet_cond, dim=(2, 3))
|
||||
if idx < indexes.shape[0]:
|
||||
feat_seq += self.task_embedding[indexes[idx][0]].to(dtype=feat_seq.dtype, device=feat_seq.device)
|
||||
|
||||
inputs.append(feat_seq.unsqueeze(1))
|
||||
condition_list.append(controlnet_cond)
|
||||
|
||||
x = torch.cat(inputs, dim=1)
|
||||
x = self.transformer_layes(x)
|
||||
|
||||
controlnet_cond_fuser = None
|
||||
for idx in range(indexes.shape[0]):
|
||||
alpha = self.spatial_ch_projs(x[:, idx])
|
||||
alpha = alpha.unsqueeze(-1).unsqueeze(-1)
|
||||
o = condition_list[idx] + alpha
|
||||
if controlnet_cond_fuser is None:
|
||||
controlnet_cond_fuser = o
|
||||
else:
|
||||
controlnet_cond_fuser += o
|
||||
return controlnet_cond_fuser
|
||||
|
||||
def forward(self, x: Tensor, hint: list[Tensor], timesteps, context, y: Tensor=None, **kwargs):
|
||||
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).to(x.dtype)
|
||||
emb = self.time_embed(t_emb)
|
||||
|
||||
guided_hint = None
|
||||
if self.control_add_embedding is not None:
|
||||
control_type = kwargs.get("control_type", None)
|
||||
|
||||
emb += self.control_add_embedding(control_type, emb.dtype, emb.device)
|
||||
if control_type is not None:
|
||||
guided_hint = self.union_controlnet_merge(hint, control_type, emb, context)
|
||||
|
||||
if guided_hint is None:
|
||||
guided_hint = self.input_hint_block(hint[0], emb, context)
|
||||
|
||||
out_output = []
|
||||
out_middle = []
|
||||
|
||||
hs = []
|
||||
if self.num_classes is not None:
|
||||
assert y.shape[0] == x.shape[0]
|
||||
emb = emb + self.label_emb(y)
|
||||
|
||||
h = x
|
||||
for module, zero_conv in zip(self.input_blocks, self.zero_convs):
|
||||
if guided_hint is not None:
|
||||
h = module(h, emb, context)
|
||||
h += guided_hint
|
||||
guided_hint = None
|
||||
else:
|
||||
h = module(h, emb, context)
|
||||
out_output.append(zero_conv(h, emb, context))
|
||||
|
||||
h = self.middle_block(h, emb, context)
|
||||
out_middle.append(self.middle_block_out(h, emb, context))
|
||||
|
||||
return {"middle": out_middle, "output": out_output}
|
||||
|
||||
|
||||
class ControlNetPlusPlusAdvanced(ControlNet, AdvancedControlBase):
|
||||
def __init__(self, control_model: ControlNetPlusPlus, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, load_device=None, manual_cast_dtype=None):
|
||||
super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype)
|
||||
AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controlnet())
|
||||
self.add_compatible_weight(ControlWeightType.CONTROLNETPLUSPLUS)
|
||||
# for IDE type hint purposes
|
||||
self.control_model: ControlNetPlusPlus
|
||||
self.cond_hint_original: Union[PlusPlusImageWrapper, PlusPlusInputGroup]
|
||||
self.cond_hint: list[Union[Tensor, None]]
|
||||
self.cond_hint_shape: Tensor = None
|
||||
self.cond_hint_types: Tensor = None
|
||||
# in case it is using the single loader
|
||||
self.single_control_type: str = None
|
||||
|
||||
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 * 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]
|
||||
if key == "input":
|
||||
raw_weights.reverse()
|
||||
return raw_weights[idx]
|
||||
return self.weights.copy_with_new_weights(new_weight_func=cn_weights_func)
|
||||
|
||||
def verify_control_type(self, model_name: str, pp_group: PlusPlusInputGroup=None):
|
||||
if pp_group is not None:
|
||||
for pp_input in pp_group.controls.values():
|
||||
if PlusPlusType.to_idx(pp_input.control_type) >= self.control_model.num_control_type:
|
||||
raise Exception(f"ControlNet++ model '{model_name}' does not support control_type '{pp_input.control_type}'.")
|
||||
if self.single_control_type is not None:
|
||||
if PlusPlusType.to_idx(self.single_control_type) >= self.control_model.num_control_type:
|
||||
raise Exception(f"ControlNet++ model '{model_name}' does not support control_type '{self.single_control_type}'.")
|
||||
|
||||
def set_cond_hint_inject(self, *args, **kwargs):
|
||||
to_return = super().set_cond_hint_inject(*args, **kwargs)
|
||||
# if not single_control_type, expect PlusPlusImageWrapper
|
||||
if self.single_control_type is None:
|
||||
# check that cond_hint is wrapped, and unwrap it
|
||||
if type(self.cond_hint_original) != PlusPlusImageWrapper:
|
||||
raise Exception("ControlNet++ (Multi) expects image input from the Load ControlNet++ Model node, NOT from anything else. Images are provided to that node via ControlNet++ Input nodes.")
|
||||
self.cond_hint_original = self.cond_hint_original.condhint.clone()
|
||||
# otherwise, expect single image input (AKA, usual controlnet input)
|
||||
else:
|
||||
# check that cond_hint is not a PlusPlusImageWrapper
|
||||
if type(self.cond_hint_original) == PlusPlusImageWrapper:
|
||||
raise Exception("ControlNet++ (Single) expects usual image input, NOT the image input from a Load ControlNet++ Model (Multi) node.")
|
||||
pp_group = PlusPlusInputGroup()
|
||||
pp_input = PlusPlusInput(self.cond_hint_original, self.single_control_type, 1.0)
|
||||
pp_group.add(pp_input)
|
||||
self.cond_hint_original = pp_group
|
||||
return to_return
|
||||
|
||||
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number, transformer_options):
|
||||
control_prev = None
|
||||
if self.previous_controlnet is not None:
|
||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number, transformer_options)
|
||||
|
||||
if self.timestep_range is not None:
|
||||
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
|
||||
if control_prev is not None:
|
||||
return control_prev
|
||||
else:
|
||||
return None
|
||||
|
||||
dtype = self.control_model.dtype
|
||||
if self.manual_cast_dtype is not None:
|
||||
dtype = self.manual_cast_dtype
|
||||
|
||||
output_dtype = x_noisy.dtype
|
||||
|
||||
# make all cond_hints appropriate dimensions
|
||||
# TODO: change this to not require cond_hint upscaling every step when self.sub_idxs is present
|
||||
if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] * self.compression_ratio != self.cond_hint_shape[2] or x_noisy.shape[3] * self.compression_ratio != self.cond_hint_shape[3]:
|
||||
if self.cond_hint is not None:
|
||||
del self.cond_hint
|
||||
self.cond_hint = [None] * self.control_model.num_control_type
|
||||
self.cond_hint_types = torch.tensor([0.0] * self.control_model.num_control_type)
|
||||
self.cond_hint_shape = None
|
||||
compression_ratio = self.compression_ratio
|
||||
# unlike normal controlnet, need to handle each input image tensor (for each type)
|
||||
for pp_type, pp_input in self.cond_hint_original.controls.items():
|
||||
pp_idx = PlusPlusType.to_idx(pp_type)
|
||||
# if negative, means no type should be selected (single only)
|
||||
if pp_idx < 0:
|
||||
pp_idx = 0
|
||||
else:
|
||||
self.cond_hint_types[pp_idx] = pp_input.strength
|
||||
# if self.cond_hint_original lengths greater or equal to latent count, subdivide
|
||||
if self.sub_idxs is not None:
|
||||
actual_cond_hint_orig = pp_input.image
|
||||
if pp_input.image.size(0) < self.full_latent_length:
|
||||
actual_cond_hint_orig = extend_to_batch_size(tensor=actual_cond_hint_orig, batch_size=self.full_latent_length)
|
||||
self.cond_hint[pp_idx] = comfy.utils.common_upscale(actual_cond_hint_orig[self.sub_idxs], x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, 'nearest-exact', "center")
|
||||
else:
|
||||
self.cond_hint[pp_idx] = comfy.utils.common_upscale(pp_input.image, x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, 'nearest-exact', "center")
|
||||
self.cond_hint[pp_idx] = self.cond_hint[pp_idx].to(device=x_noisy.device, dtype=dtype)
|
||||
self.cond_hint_shape = self.cond_hint[pp_idx].shape
|
||||
# prepare cond_hint_controls to match batchsize
|
||||
if self.cond_hint_types.count_nonzero() == 0:
|
||||
self.cond_hint_types = None
|
||||
else:
|
||||
self.cond_hint_types = self.cond_hint_types.unsqueeze(0).to(device=x_noisy.device, dtype=dtype).repeat(x_noisy.shape[0], 1)
|
||||
for i in range(len(self.cond_hint)):
|
||||
if self.cond_hint[i] is not None:
|
||||
if x_noisy.shape[0] != self.cond_hint[i].shape[0]:
|
||||
self.cond_hint[i] = broadcast_image_to_extend(self.cond_hint[i], x_noisy.shape[0], batched_number)
|
||||
if self.cond_hint_types is not None and x_noisy.shape[0] != self.cond_hint_types.shape[0]:
|
||||
self.cond_hint_types = broadcast_image_to_extend(self.cond_hint_types, x_noisy.shape[0], batched_number, False)
|
||||
|
||||
# prepare mask_cond_hint
|
||||
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, dtype=dtype)
|
||||
|
||||
context = cond.get('crossattn_controlnet', cond['c_crossattn'])
|
||||
y = cond.get('y', None)
|
||||
if y is not None:
|
||||
y = comfy.model_base.convert_tensor(y, dtype, x_noisy.device)
|
||||
timestep = self.model_sampling_current.timestep(t)
|
||||
x_noisy = self.model_sampling_current.calculate_input(t, x_noisy)
|
||||
|
||||
control = self.control_model(x=x_noisy.to(dtype), hint=self.cond_hint, timesteps=timestep.float(), context=comfy.model_management.cast_to_device(context, x_noisy.device, dtype), y=y, control_type=self.cond_hint_types)
|
||||
return self.control_merge(control, control_prev, output_dtype)
|
||||
|
||||
def copy(self):
|
||||
c = ControlNetPlusPlusAdvanced(self.control_model, self.timestep_keyframes, global_average_pooling=self.global_average_pooling, load_device=self.load_device, manual_cast_dtype=self.manual_cast_dtype)
|
||||
self.copy_to(c)
|
||||
self.copy_to_advanced(c)
|
||||
c.single_control_type = self.single_control_type
|
||||
return c
|
||||
|
||||
|
||||
def load_controlnetplusplus(ckpt_path: str, timestep_keyframe: TimestepKeyframeGroup=None, model=None):
|
||||
controlnet_data = comfy.utils.load_torch_file(ckpt_path, safe_load=True)
|
||||
# check that actually is ControlNet++ model
|
||||
if "task_embedding" not in controlnet_data:
|
||||
raise Exception(f"'{ckpt_path}' is not a valid ControlNet++ model.")
|
||||
|
||||
controlnet_config = None
|
||||
supported_inference_dtypes = None
|
||||
|
||||
if "controlnet_cond_embedding.conv_in.weight" in controlnet_data: #diffusers format
|
||||
controlnet_config = comfy.model_detection.unet_config_from_diffusers_unet(controlnet_data)
|
||||
diffusers_keys = comfy.utils.unet_to_diffusers(controlnet_config)
|
||||
diffusers_keys["controlnet_mid_block.weight"] = "middle_block_out.0.weight"
|
||||
diffusers_keys["controlnet_mid_block.bias"] = "middle_block_out.0.bias"
|
||||
|
||||
count = 0
|
||||
loop = True
|
||||
while loop:
|
||||
suffix = [".weight", ".bias"]
|
||||
for s in suffix:
|
||||
k_in = "controlnet_down_blocks.{}{}".format(count, s)
|
||||
k_out = "zero_convs.{}.0{}".format(count, s)
|
||||
if k_in not in controlnet_data:
|
||||
loop = False
|
||||
break
|
||||
diffusers_keys[k_in] = k_out
|
||||
count += 1
|
||||
|
||||
count = 0
|
||||
loop = True
|
||||
while loop:
|
||||
suffix = [".weight", ".bias"]
|
||||
for s in suffix:
|
||||
if count == 0:
|
||||
k_in = "controlnet_cond_embedding.conv_in{}".format(s)
|
||||
else:
|
||||
k_in = "controlnet_cond_embedding.blocks.{}{}".format(count - 1, s)
|
||||
k_out = "input_hint_block.{}{}".format(count * 2, s)
|
||||
if k_in not in controlnet_data:
|
||||
k_in = "controlnet_cond_embedding.conv_out{}".format(s)
|
||||
loop = False
|
||||
diffusers_keys[k_in] = k_out
|
||||
count += 1
|
||||
|
||||
new_sd = {}
|
||||
for k in diffusers_keys:
|
||||
if k in controlnet_data:
|
||||
new_sd[diffusers_keys[k]] = controlnet_data.pop(k)
|
||||
|
||||
if "control_add_embedding.linear_1.bias" in controlnet_data: #Union Controlnet
|
||||
controlnet_config["union_controlnet_num_control_type"] = controlnet_data["task_embedding"].shape[0]
|
||||
for k in list(controlnet_data.keys()):
|
||||
new_k = k.replace('.attn.in_proj_', '.attn.in_proj.')
|
||||
new_sd[new_k] = controlnet_data.pop(k)
|
||||
|
||||
leftover_keys = controlnet_data.keys()
|
||||
if len(leftover_keys) > 0:
|
||||
logger.warning("leftover ControlNet++ keys: {}".format(leftover_keys))
|
||||
controlnet_data = new_sd
|
||||
elif "controlnet_blocks.0.weight" in controlnet_data: #SD3 diffusers format
|
||||
raise Exception("Unexpected SD3 diffusers format for ControlNet++ model. Something is very wrong.")
|
||||
|
||||
pth_key = 'control_model.zero_convs.0.0.weight'
|
||||
pth = False
|
||||
key = 'zero_convs.0.0.weight'
|
||||
if pth_key in controlnet_data:
|
||||
pth = True
|
||||
key = pth_key
|
||||
prefix = "control_model."
|
||||
elif key in controlnet_data:
|
||||
prefix = ""
|
||||
else:
|
||||
raise Exception("Unexpected T2IAdapter format for ControlNet++ model. Something is very wrong.")
|
||||
|
||||
if controlnet_config is None:
|
||||
model_config = comfy.model_detection.model_config_from_unet(controlnet_data, prefix, True)
|
||||
supported_inference_dtypes = model_config.supported_inference_dtypes
|
||||
controlnet_config = model_config.unet_config
|
||||
|
||||
load_device = comfy.model_management.get_torch_device()
|
||||
if supported_inference_dtypes is None:
|
||||
unet_dtype = comfy.model_management.unet_dtype()
|
||||
else:
|
||||
unet_dtype = comfy.model_management.unet_dtype(supported_dtypes=supported_inference_dtypes)
|
||||
|
||||
manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device)
|
||||
if manual_cast_dtype is not None:
|
||||
controlnet_config["operations"] = comfy.ops.manual_cast
|
||||
controlnet_config["dtype"] = unet_dtype
|
||||
controlnet_config.pop("out_channels")
|
||||
controlnet_config["hint_channels"] = controlnet_data["{}input_hint_block.0.weight".format(prefix)].shape[1]
|
||||
control_model = ControlNetPlusPlus(**controlnet_config)
|
||||
|
||||
if pth:
|
||||
if 'difference' in controlnet_data:
|
||||
if model is not None:
|
||||
comfy.model_management.load_models_gpu([model])
|
||||
model_sd = model.model_state_dict()
|
||||
for x in controlnet_data:
|
||||
c_m = "control_model."
|
||||
if x.startswith(c_m):
|
||||
sd_key = "diffusion_model.{}".format(x[len(c_m):])
|
||||
if sd_key in model_sd:
|
||||
cd = controlnet_data[x]
|
||||
cd += model_sd[sd_key].type(cd.dtype).to(cd.device)
|
||||
else:
|
||||
logger.warning("WARNING: Loaded a diff controlnet without a model. It will very likely not work.")
|
||||
|
||||
class WeightsLoader(torch.nn.Module):
|
||||
pass
|
||||
w = WeightsLoader()
|
||||
w.control_model = control_model
|
||||
missing, unexpected = w.load_state_dict(controlnet_data, strict=False)
|
||||
else:
|
||||
missing, unexpected = control_model.load_state_dict(controlnet_data, strict=False)
|
||||
|
||||
if len(missing) > 0:
|
||||
logger.warning("missing ControlNet++ keys: {}".format(missing))
|
||||
|
||||
if len(unexpected) > 0:
|
||||
logger.debug("unexpected ControlNet++ keys: {}".format(unexpected))
|
||||
|
||||
global_average_pooling = False
|
||||
filename = os.path.splitext(ckpt_path)[0]
|
||||
if filename.endswith("_shuffle") or filename.endswith("_shuffle_fp16"): #TODO: smarter way of enabling global_average_pooling
|
||||
global_average_pooling = True
|
||||
|
||||
control = ControlNetPlusPlusAdvanced(control_model, timestep_keyframes=timestep_keyframe, global_average_pooling=global_average_pooling, load_device=load_device, manual_cast_dtype=manual_cast_dtype)
|
||||
return control
|
||||
+300
-673
File diff suppressed because it is too large
Load Diff
+935
-153
File diff suppressed because it is too large
Load Diff
@@ -311,8 +311,7 @@ class SVDControlNet(nn.Module):
|
||||
|
||||
guided_hint = self.input_hint_block(hint, emb, context, time_context=time_context, num_video_frames=num_video_frames, image_only_indicator=image_only_indicator)
|
||||
|
||||
out_output = []
|
||||
out_middle = []
|
||||
outs = []
|
||||
|
||||
hs = []
|
||||
if self.num_classes is not None:
|
||||
@@ -327,12 +326,12 @@ class SVDControlNet(nn.Module):
|
||||
guided_hint = None
|
||||
else:
|
||||
h = module(h, emb, context, time_context=time_context, num_video_frames=num_video_frames, image_only_indicator=image_only_indicator)
|
||||
out_output.append(zero_conv(h, emb, context, time_context=time_context, num_video_frames=num_video_frames, image_only_indicator=image_only_indicator))
|
||||
outs.append(zero_conv(h, emb, context, time_context=time_context, num_video_frames=num_video_frames, image_only_indicator=image_only_indicator))
|
||||
|
||||
h = self.middle_block(h, emb, context, time_context=time_context, num_video_frames=num_video_frames, image_only_indicator=image_only_indicator)
|
||||
out_middle.append(self.middle_block_out(h, emb, context, time_context=time_context, num_video_frames=num_video_frames, image_only_indicator=image_only_indicator))
|
||||
outs.append(self.middle_block_out(h, emb, context, time_context=time_context, num_video_frames=num_video_frames, image_only_indicator=image_only_indicator))
|
||||
|
||||
return {"middle": out_middle, "output": out_output}
|
||||
return outs
|
||||
|
||||
|
||||
TEMPORAL_TRANSFORMER_BLOCKS = {
|
||||
|
||||
@@ -1,112 +0,0 @@
|
||||
####################################################################################################
|
||||
# DinkLink is my method of sharing classes/functions between my nodes.
|
||||
#
|
||||
# My DinkLink-compatible nodes will inject comfy.hooks with a __DINKLINK attr
|
||||
# that stores a dictionary, where any of my node packs can store their stuff.
|
||||
#
|
||||
# It is not intended to be accessed by node packs that I don't develop, so things may change
|
||||
# at any time.
|
||||
#
|
||||
# DinkLink also serves as a proof-of-concept for a future ComfyUI implementation of
|
||||
# purposely exposing node pack classes/functions with other node packs.
|
||||
####################################################################################################
|
||||
from __future__ import annotations
|
||||
from typing import Union
|
||||
from torch import Tensor, nn
|
||||
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
import comfy.hooks
|
||||
|
||||
DINKLINK = "__DINKLINK"
|
||||
|
||||
|
||||
def init_dinklink():
|
||||
create_dinklink()
|
||||
prepare_dinklink()
|
||||
|
||||
def create_dinklink():
|
||||
if not hasattr(comfy.hooks, DINKLINK):
|
||||
setattr(comfy.hooks, DINKLINK, {})
|
||||
|
||||
def get_dinklink() -> dict[str, dict[str]]:
|
||||
create_dinklink()
|
||||
return getattr(comfy.hooks, DINKLINK)
|
||||
|
||||
|
||||
class DinkLinkConst:
|
||||
VERSION = "version"
|
||||
# ADE
|
||||
ADE = "ADE"
|
||||
ADE_ANIMATEDIFFMODEL = "AnimateDiffModel"
|
||||
ADE_ANIMATEDIFFINFO = "AnimateDiffInfo"
|
||||
ADE_CREATE_MOTIONMODELPATCHER = "create_MotionModelPatcher"
|
||||
|
||||
def prepare_dinklink():
|
||||
pass
|
||||
|
||||
|
||||
class InterfaceAnimateDiffInfo:
|
||||
'''Class only used for IDE type hints; interface of ADE's AnimateDiffInfo'''
|
||||
def __init__(self, sd_type: str, mm_format: str, mm_version: str, mm_name: str):
|
||||
self.sd_type = sd_type
|
||||
self.mm_format = mm_format
|
||||
self.mm_version = mm_version
|
||||
self.mm_name = mm_name
|
||||
|
||||
|
||||
class InterfaceAnimateDiffModel(nn.Module):
|
||||
'''Class only used for IDE type hints; interface of ADE's AnimateDiffModel'''
|
||||
def __init__(self, mm_state_dict: dict[str, Tensor], mm_info: InterfaceAnimateDiffInfo, init_kwargs: dict[str]={}):
|
||||
pass
|
||||
|
||||
def set_video_length(self, video_length: int, full_length: int) -> None:
|
||||
raise NotImplemented()
|
||||
|
||||
def set_scale(self, scale: Union[float, Tensor, None], per_block_list: Union[list, None]=None) -> None:
|
||||
raise NotImplemented()
|
||||
|
||||
def set_effect(self, multival: Union[float, Tensor, None], per_block_list: Union[list, None]=None) -> None:
|
||||
raise NotImplemented()
|
||||
|
||||
def cleanup(self):
|
||||
raise NotImplemented()
|
||||
|
||||
def inject(self, model: ModelPatcher):
|
||||
pass
|
||||
|
||||
def eject(self, model: ModelPatcher):
|
||||
pass
|
||||
|
||||
|
||||
def get_CreateMotionModelPatcher(throw_exception=True):
|
||||
d = get_dinklink()
|
||||
try:
|
||||
link_ade = d[DinkLinkConst.ADE]
|
||||
return link_ade[DinkLinkConst.ADE_CREATE_MOTIONMODELPATCHER]
|
||||
except KeyError:
|
||||
if throw_exception:
|
||||
raise Exception("Could not get create_MotionModelPatcher function. AnimateDiff-Evolved nodes need to be installed to use SparseCtrl; " + \
|
||||
"they are either not installed or are of an insufficient version.")
|
||||
return None
|
||||
|
||||
def get_AnimateDiffModel(throw_exception=True):
|
||||
d = get_dinklink()
|
||||
try:
|
||||
link_ade = d[DinkLinkConst.ADE]
|
||||
return link_ade[DinkLinkConst.ADE_ANIMATEDIFFMODEL]
|
||||
except KeyError:
|
||||
if throw_exception:
|
||||
raise Exception("Could not get AnimateDiffModel class. AnimateDiff-Evolved nodes need to be installed to use SparseCtrl; " + \
|
||||
"they are either not installed or are of an insufficient version.")
|
||||
return None
|
||||
|
||||
def get_AnimateDiffInfo(throw_exception=True) -> InterfaceAnimateDiffInfo:
|
||||
d = get_dinklink()
|
||||
try:
|
||||
link_ade = d[DinkLinkConst.ADE]
|
||||
return link_ade[DinkLinkConst.ADE_ANIMATEDIFFINFO]
|
||||
except KeyError:
|
||||
if throw_exception:
|
||||
raise Exception("Could not get AnimateDiffInfo class - AnimateDiff-Evolved nodes need to be installed to use SparseCtrl; " + \
|
||||
"they are either not installed or are of an insufficient version.")
|
||||
return None
|
||||
@@ -1,47 +0,0 @@
|
||||
from .logger import logger
|
||||
|
||||
def image(src):
|
||||
return f'<img src={src} style="width: 0px; min-width: 100%">'
|
||||
def video(src):
|
||||
return f'<video src={src} autoplay muted loop controls controlslist="nodownload noremoteplayback noplaybackrate" style="width: 0px; min-width: 100%" class="VHS_loopedvideo">'
|
||||
def short_desc(desc):
|
||||
return f'<div id=VHS_shortdesc style="font-size: .8em">{desc}</div>'
|
||||
|
||||
descriptions = {
|
||||
}
|
||||
|
||||
sizes = ['1.4','1.2','1']
|
||||
def as_html(entry, depth=0):
|
||||
if isinstance(entry, dict):
|
||||
size = 0.8 if depth < 2 else 1
|
||||
html = ''
|
||||
for k in entry:
|
||||
if k == "collapsed":
|
||||
continue
|
||||
collapse_single = k.endswith("_collapsed")
|
||||
if collapse_single:
|
||||
name = k[:-len("_collapsed")]
|
||||
else:
|
||||
name = k
|
||||
collapse_flag = ' VHS_precollapse' if entry.get("collapsed", False) or collapse_single else ''
|
||||
html += f'<div vhs_title=\"{name}\" style=\"display: flex; font-size: {size}em\" class=\"VHS_collapse{collapse_flag}\"><div style=\"color: #AAA; height: 1.5em;\">[<span style=\"font-family: monospace\">-</span>]</div><div style=\"width: 100%\">{name}: {as_html(entry[k], depth=depth+1)}</div></div>'
|
||||
return html
|
||||
if isinstance(entry, list):
|
||||
html = ''
|
||||
for i in entry:
|
||||
html += f'<div>{as_html(i, depth=depth)}</div>'
|
||||
return html
|
||||
return str(entry)
|
||||
|
||||
def format_descriptions(nodes):
|
||||
for k in descriptions:
|
||||
if k.endswith("_collapsed"):
|
||||
k = k[:-len("_collapsed")]
|
||||
nodes[k].DESCRIPTION = as_html(descriptions[k])
|
||||
# undocumented_nodes = []
|
||||
# for k in nodes:
|
||||
# if not hasattr(nodes[k], "DESCRIPTION"):
|
||||
# undocumented_nodes.append(k)
|
||||
# if len(undocumented_nodes) > 0:
|
||||
# logger.info(f"Undocumented nodes: {undocumented_nodes}")
|
||||
|
||||
+170
-71
@@ -1,25 +1,162 @@
|
||||
import comfy.sample
|
||||
import numpy as np
|
||||
from torch import Tensor
|
||||
|
||||
from .nodes_main import (ControlNetLoaderAdvanced, DiffControlNetLoaderAdvanced,
|
||||
AdvancedControlNetApply, AdvancedControlNetApplySingle)
|
||||
from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights,
|
||||
SoftControlNetWeightsSD15, CustomControlNetWeightsSD15, CustomControlNetWeightsFlux,
|
||||
SoftT2IAdapterWeights, CustomT2IAdapterWeights, ExtrasMiddleMultNode)
|
||||
import folder_paths
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .control import load_controlnet, convert_to_advanced, is_advanced_controlnet
|
||||
from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, BIGMAX
|
||||
from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, ScaledSoftUniversalWeights, SoftControlNetWeights, CustomControlNetWeights,
|
||||
SoftT2IAdapterWeights, CustomT2IAdapterWeights)
|
||||
from .nodes_keyframes import (LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode,
|
||||
TimestepKeyframeNode, TimestepKeyframeInterpolationNode, TimestepKeyframeFromStrengthListNode)
|
||||
from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor, SparseWeightExtras
|
||||
from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor
|
||||
from .nodes_reference import ReferenceControlNetNode, ReferenceControlFinetune, ReferencePreprocessorNode
|
||||
from .nodes_plusplus import PlusPlusLoaderAdvanced, PlusPlusLoaderSingle, PlusPlusInputNode
|
||||
from .nodes_ctrlora import CtrLoRALoader
|
||||
from .nodes_loosecontrol import ControlNetLoaderWithLoraAdvanced
|
||||
from .nodes_deprecated import (LoadImagesFromDirectory, ScaledSoftUniversalWeightsDeprecated,
|
||||
SoftControlNetWeightsDeprecated, CustomControlNetWeightsDeprecated,
|
||||
SoftT2IAdapterWeightsDeprecated, CustomT2IAdapterWeightsDeprecated,
|
||||
AdvancedControlNetApplyDEPR, AdvancedControlNetApplySingleDEPR,
|
||||
ControlNetLoaderAdvancedDEPR, DiffControlNetLoaderAdvancedDEPR)
|
||||
from .nodes_deprecated import LoadImagesFromDirectory
|
||||
from .logger import logger
|
||||
|
||||
|
||||
class ControlNetLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"control_net_name": (folder_paths.get_filename_list("controlnet"), ),
|
||||
},
|
||||
"optional": {
|
||||
"timestep_keyframe": ("TIMESTEP_KEYFRAME", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝"
|
||||
|
||||
def load_controlnet(self, control_net_name,
|
||||
timestep_keyframe: TimestepKeyframeGroup=None
|
||||
):
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
controlnet = load_controlnet(controlnet_path, timestep_keyframe)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class DiffControlNetLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"control_net_name": (folder_paths.get_filename_list("controlnet"), )
|
||||
},
|
||||
"optional": {
|
||||
"timestep_keyframe": ("TIMESTEP_KEYFRAME", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝"
|
||||
|
||||
def load_controlnet(self, control_net_name, model,
|
||||
timestep_keyframe: TimestepKeyframeGroup=None
|
||||
):
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
controlnet = load_controlnet(controlnet_path, timestep_keyframe, model)
|
||||
if is_advanced_controlnet(controlnet):
|
||||
controlnet.verify_all_weights()
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class AdvancedControlNetApply:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"control_net": ("CONTROL_NET", ),
|
||||
"image": ("IMAGE", ),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
|
||||
},
|
||||
"optional": {
|
||||
"mask_optional": ("MASK", ),
|
||||
"timestep_kf": ("TIMESTEP_KEYFRAME", ),
|
||||
"latent_kf_override": ("LATENT_KEYFRAME", ),
|
||||
"weights_override": ("CONTROL_NET_WEIGHTS", ),
|
||||
"model_optional": ("MODEL",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING","CONDITIONING","MODEL",)
|
||||
RETURN_NAMES = ("positive", "negative", "model_opt")
|
||||
FUNCTION = "apply_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝"
|
||||
|
||||
def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent,
|
||||
mask_optional: Tensor=None, model_optional: ModelPatcher=None,
|
||||
timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None,
|
||||
weights_override: ControlWeights=None):
|
||||
if strength == 0:
|
||||
return (positive, negative, model_optional)
|
||||
if model_optional:
|
||||
model_optional = model_optional.clone()
|
||||
|
||||
control_hint = image.movedim(-1,1)
|
||||
cnets = {}
|
||||
|
||||
out = []
|
||||
for conditioning in [positive, negative]:
|
||||
c = []
|
||||
for t in conditioning:
|
||||
d = t[1].copy()
|
||||
|
||||
prev_cnet = d.get('control', None)
|
||||
if prev_cnet in cnets:
|
||||
c_net = cnets[prev_cnet]
|
||||
else:
|
||||
# copy, convert to advanced if needed, and set cond
|
||||
c_net = convert_to_advanced(control_net.copy()).set_cond_hint(control_hint, strength, (start_percent, end_percent))
|
||||
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)
|
||||
# apply optional parameters and overrides, if provided
|
||||
if timestep_kf is not None:
|
||||
c_net.set_timestep_keyframes(timestep_kf)
|
||||
if latent_kf_override is not None:
|
||||
c_net.latent_keyframe_override = latent_kf_override
|
||||
if weights_override is not None:
|
||||
c_net.weights_override = weights_override
|
||||
# verify weights are compatible
|
||||
c_net.verify_all_weights()
|
||||
# set cond hint mask
|
||||
if mask_optional is not None:
|
||||
mask_optional = mask_optional.clone()
|
||||
# if not in the form of a batch, make it so
|
||||
if len(mask_optional.shape) < 3:
|
||||
mask_optional = mask_optional.unsqueeze(0)
|
||||
c_net.set_cond_hint_mask(mask_optional)
|
||||
c_net.set_previous_controlnet(prev_cnet)
|
||||
cnets[prev_cnet] = c_net
|
||||
|
||||
d['control'] = c_net
|
||||
d['control_apply_to_uncond'] = False
|
||||
n = [t[0], d]
|
||||
c.append(n)
|
||||
out.append(c)
|
||||
return (out[0], out[1], model_optional)
|
||||
|
||||
|
||||
# NODE MAPPING
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
# Keyframes
|
||||
@@ -31,34 +168,24 @@ NODE_CLASS_MAPPINGS = {
|
||||
"LatentKeyframeBatchedGroup": LatentKeyframeBatchedGroupNode,
|
||||
"LatentKeyframeGroup": LatentKeyframeGroupNode,
|
||||
# Conditioning
|
||||
"ACN_AdvancedControlNetApply_v2": AdvancedControlNetApply,
|
||||
"ACN_AdvancedControlNetApplySingle_v2": AdvancedControlNetApplySingle,
|
||||
"ACN_AdvancedControlNetApply": AdvancedControlNetApply,
|
||||
# Loaders
|
||||
"ACN_ControlNetLoaderAdvanced": ControlNetLoaderAdvanced,
|
||||
"ACN_DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced,
|
||||
"ControlNetLoaderAdvanced": ControlNetLoaderAdvanced,
|
||||
"DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvanced,
|
||||
# Weights
|
||||
"ACN_ScaledSoftControlNetWeights": ScaledSoftUniversalWeights,
|
||||
"ScaledSoftControlNetWeights": ScaledSoftUniversalWeights,
|
||||
"ScaledSoftMaskedUniversalWeights": ScaledSoftMaskedUniversalWeights,
|
||||
"ACN_SoftControlNetWeightsSD15": SoftControlNetWeightsSD15,
|
||||
"ACN_CustomControlNetWeightsSD15": CustomControlNetWeightsSD15,
|
||||
"ACN_CustomControlNetWeightsFlux": CustomControlNetWeightsFlux,
|
||||
"ACN_SoftT2IAdapterWeights": SoftT2IAdapterWeights,
|
||||
"ACN_CustomT2IAdapterWeights": CustomT2IAdapterWeights,
|
||||
"SoftControlNetWeights": SoftControlNetWeights,
|
||||
"CustomControlNetWeights": CustomControlNetWeights,
|
||||
"SoftT2IAdapterWeights": SoftT2IAdapterWeights,
|
||||
"CustomT2IAdapterWeights": CustomT2IAdapterWeights,
|
||||
"ACN_DefaultUniversalWeights": DefaultWeights,
|
||||
"ACN_ExtrasMiddleMult": ExtrasMiddleMultNode,
|
||||
# SparseCtrl
|
||||
"ACN_SparseCtrlRGBPreprocessor": RgbSparseCtrlPreprocessor,
|
||||
"ACN_SparseCtrlLoaderAdvanced": SparseCtrlLoaderAdvanced,
|
||||
"ACN_SparseCtrlMergedLoaderAdvanced": SparseCtrlMergedLoaderAdvanced,
|
||||
"ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode,
|
||||
"ACN_SparseCtrlSpreadMethodNode": SparseSpreadMethodNode,
|
||||
"ACN_SparseCtrlWeightExtras": SparseWeightExtras,
|
||||
# ControlNet++
|
||||
"ACN_ControlNet++LoaderSingle": PlusPlusLoaderSingle,
|
||||
"ACN_ControlNet++LoaderAdvanced": PlusPlusLoaderAdvanced,
|
||||
"ACN_ControlNet++InputNode": PlusPlusInputNode,
|
||||
# CtrLoRA
|
||||
"ACN_CtrLoRALoader": CtrLoRALoader,
|
||||
# Reference
|
||||
"ACN_ReferencePreprocessor": ReferencePreprocessorNode,
|
||||
"ACN_ReferenceControlNet": ReferenceControlNetNode,
|
||||
@@ -67,55 +194,36 @@ NODE_CLASS_MAPPINGS = {
|
||||
#"ACN_ControlNetLoaderWithLoraAdvanced": ControlNetLoaderWithLoraAdvanced,
|
||||
# Deprecated
|
||||
"LoadImagesFromDirectory": LoadImagesFromDirectory,
|
||||
"ScaledSoftControlNetWeights": ScaledSoftUniversalWeightsDeprecated,
|
||||
"SoftControlNetWeights": SoftControlNetWeightsDeprecated,
|
||||
"CustomControlNetWeights": CustomControlNetWeightsDeprecated,
|
||||
"SoftT2IAdapterWeights": SoftT2IAdapterWeightsDeprecated,
|
||||
"CustomT2IAdapterWeights": CustomT2IAdapterWeightsDeprecated,
|
||||
"ACN_AdvancedControlNetApply": AdvancedControlNetApplyDEPR,
|
||||
"ACN_AdvancedControlNetApplySingle": AdvancedControlNetApplySingleDEPR,
|
||||
"ControlNetLoaderAdvanced": ControlNetLoaderAdvancedDEPR,
|
||||
"DiffControlNetLoaderAdvanced": DiffControlNetLoaderAdvancedDEPR,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Keyframes
|
||||
"TimestepKeyframe": "Timestep Keyframe 🛂🅐🅒🅝",
|
||||
"ACN_TimestepKeyframeInterpolation": "Timestep Keyframe Interp. 🛂🅐🅒🅝",
|
||||
"ACN_TimestepKeyframeInterpolation": "Timestep Keyframe Interpolation 🛂🅐🅒🅝",
|
||||
"ACN_TimestepKeyframeFromStrengthList": "Timestep Keyframe From List 🛂🅐🅒🅝",
|
||||
"LatentKeyframe": "Latent Keyframe 🛂🅐🅒🅝",
|
||||
"LatentKeyframeTiming": "Latent Keyframe Interp. 🛂🅐🅒🅝",
|
||||
"LatentKeyframeTiming": "Latent Keyframe Interpolation 🛂🅐🅒🅝",
|
||||
"LatentKeyframeBatchedGroup": "Latent Keyframe From List 🛂🅐🅒🅝",
|
||||
"LatentKeyframeGroup": "Latent Keyframe Group 🛂🅐🅒🅝",
|
||||
# Conditioning
|
||||
"ACN_AdvancedControlNetApply_v2": "Apply Advanced ControlNet 🛂🅐🅒🅝",
|
||||
"ACN_AdvancedControlNetApplySingle_v2": "Apply Advanced ControlNet(1) 🛂🅐🅒🅝",
|
||||
"ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝",
|
||||
# Loaders
|
||||
"ACN_ControlNetLoaderAdvanced": "Load Advanced ControlNet Model 🛂🅐🅒🅝",
|
||||
"ACN_DiffControlNetLoaderAdvanced": "Load Advanced ControlNet Model (diff) 🛂🅐🅒🅝",
|
||||
"ControlNetLoaderAdvanced": "Load Advanced ControlNet Model 🛂🅐🅒🅝",
|
||||
"DiffControlNetLoaderAdvanced": "Load Advanced ControlNet Model (diff) 🛂🅐🅒🅝",
|
||||
# Weights
|
||||
"ACN_ScaledSoftControlNetWeights": "Scaled Soft Weights 🛂🅐🅒🅝",
|
||||
"ScaledSoftControlNetWeights": "Scaled Soft Weights 🛂🅐🅒🅝",
|
||||
"ScaledSoftMaskedUniversalWeights": "Scaled Soft Masked Weights 🛂🅐🅒🅝",
|
||||
"ACN_SoftControlNetWeightsSD15": "ControlNet Soft Weights [SD1.5] 🛂🅐🅒🅝",
|
||||
"ACN_CustomControlNetWeightsSD15": "ControlNet Custom Weights [SD1.5] 🛂🅐🅒🅝",
|
||||
"ACN_CustomControlNetWeightsFlux": "ControlNet Custom Weights [Flux] 🛂🅐🅒🅝",
|
||||
"ACN_SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝",
|
||||
"ACN_CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝",
|
||||
"ACN_DefaultUniversalWeights": "Default Weights 🛂🅐🅒🅝",
|
||||
"ACN_ExtrasMiddleMult": "Middle Weight Extras 🛂🅐🅒🅝",
|
||||
"SoftControlNetWeights": "ControlNet Soft Weights 🛂🅐🅒🅝",
|
||||
"CustomControlNetWeights": "ControlNet Custom Weights 🛂🅐🅒🅝",
|
||||
"SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝",
|
||||
"CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝",
|
||||
"ACN_DefaultUniversalWeights": "Force Default Weights 🛂🅐🅒🅝",
|
||||
# SparseCtrl
|
||||
"ACN_SparseCtrlRGBPreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlMergedLoaderAdvanced": "🧪Load Merged SparseCtrl Model 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlSpreadMethodNode": "SparseCtrl Spread Method 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlWeightExtras": "SparseCtrl Weight Extras 🛂🅐🅒🅝",
|
||||
# ControlNet++
|
||||
"ACN_ControlNet++LoaderSingle": "Load ControlNet++ Model (Single) 🛂🅐🅒🅝",
|
||||
"ACN_ControlNet++LoaderAdvanced": "Load ControlNet++ Model (Multi) 🛂🅐🅒🅝",
|
||||
"ACN_ControlNet++InputNode": "ControlNet++ Input 🛂🅐🅒🅝",
|
||||
# CtrLoRA
|
||||
"ACN_CtrLoRALoader": "Load CtrLoRA Model 🛂🅐🅒🅝",
|
||||
# Reference
|
||||
"ACN_ReferencePreprocessor": "Reference Preproccessor 🛂🅐🅒🅝",
|
||||
"ACN_ReferenceControlNet": "Reference ControlNet 🛂🅐🅒🅝",
|
||||
@@ -124,13 +232,4 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
#"ACN_ControlNetLoaderWithLoraAdvanced": "Load Adv. ControlNet Model w/ LoRA 🛂🅐🅒🅝",
|
||||
# Deprecated
|
||||
"LoadImagesFromDirectory": "🚫Load Images [DEPRECATED] 🛂🅐🅒🅝",
|
||||
"ScaledSoftControlNetWeights": "Scaled Soft Weights 🛂🅐🅒🅝",
|
||||
"SoftControlNetWeights": "ControlNet Soft Weights 🛂🅐🅒🅝",
|
||||
"CustomControlNetWeights": "ControlNet Custom Weights 🛂🅐🅒🅝",
|
||||
"SoftT2IAdapterWeights": "T2IAdapter Soft Weights 🛂🅐🅒🅝",
|
||||
"CustomT2IAdapterWeights": "T2IAdapter Custom Weights 🛂🅐🅒🅝",
|
||||
"ACN_AdvancedControlNetApply": "Apply Advanced ControlNet 🛂🅐🅒🅝",
|
||||
"ACN_AdvancedControlNetApplySingle": "Apply Advanced ControlNet(1) 🛂🅐🅒🅝",
|
||||
"ControlNetLoaderAdvanced": "Load Advanced ControlNet Model 🛂🅐🅒🅝",
|
||||
"DiffControlNetLoaderAdvanced": "Load Advanced ControlNet Model (diff) 🛂🅐🅒🅝",
|
||||
}
|
||||
|
||||
@@ -1,25 +0,0 @@
|
||||
import folder_paths
|
||||
|
||||
from .control_ctrlora import load_ctrlora
|
||||
|
||||
|
||||
class CtrLoRALoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"base": (folder_paths.get_filename_list("controlnet"), ),
|
||||
"lora": (folder_paths.get_filename_list("controlnet"), ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET",)
|
||||
FUNCTION = "load_controlnet_plusplus"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/CtrLoRA"
|
||||
|
||||
def load_controlnet_plusplus(self, base: str, lora: str):
|
||||
base_path = folder_paths.get_full_path("controlnet", base)
|
||||
lora_path = folder_paths.get_full_path("controlnet", lora)
|
||||
controlnet = load_ctrlora(base_path, lora_path)
|
||||
return (controlnet,)
|
||||
@@ -1,13 +1,10 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
import folder_paths
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image, ImageOps
|
||||
from .control import load_controlnet, is_advanced_controlnet
|
||||
from .nodes_main import AdvancedControlNetApply
|
||||
from .utils import BIGMAX, ControlWeights, TimestepKeyframeGroup, TimestepKeyframe, get_properly_arranged_t2i_weights
|
||||
from .utils import BIGMAX
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -72,345 +69,3 @@ class LoadImagesFromDirectory:
|
||||
raise FileNotFoundError(f"No images could be loaded from directory '{directory}'.")
|
||||
|
||||
return (torch.cat(images, dim=0), torch.stack(masks, dim=0), image_count)
|
||||
|
||||
|
||||
class ScaledSoftUniversalWeightsDeprecated:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
RETURN_NAMES = ("CN_WEIGHTS", "TK_SHORTCUT")
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def load_weights(self, base_multiplier, flip_weights, uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights = ControlWeights.universal(base_multiplier=base_multiplier, uncond_multiplier=uncond_multiplier, extras=cn_extras)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class SoftControlNetWeightsDeprecated:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 0.09941396206337118, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 0.12050177219802567, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 0.14606275417942507, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 0.17704576264172736, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_04": ("FLOAT", {"default": 0.214600924414215, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_05": ("FLOAT", {"default": 0.26012233262329093, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_06": ("FLOAT", {"default": 0.3152997971191405, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_07": ("FLOAT", {"default": 0.3821815722656249, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_08": ("FLOAT", {"default": 0.4632503906249999, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_09": ("FLOAT", {"default": 0.561515625, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_10": ("FLOAT", {"default": 0.6806249999999999, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_11": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
DEPRECATED = True
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
RETURN_NAMES = ("CN_WEIGHTS", "TK_SHORTCUT")
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights_output = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11]
|
||||
weights_middle = [weight_12]
|
||||
weights = ControlWeights.controlnet(weights_output=weights_output, weights_middle=weights_middle, uncond_multiplier=uncond_multiplier, extras=cn_extras)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class CustomControlNetWeightsDeprecated:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_04": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_05": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_06": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_07": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_08": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_09": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_10": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_11": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
DEPRECATED = True
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
RETURN_NAMES = ("CN_WEIGHTS", "TK_SHORTCUT")
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights_output = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11]
|
||||
weights_middle = [weight_12]
|
||||
weights = ControlWeights.controlnet(weights_output=weights_output, weights_middle=weights_middle, uncond_multiplier=uncond_multiplier, extras=cn_extras)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class SoftT2IAdapterWeightsDeprecated:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 0.62, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
DEPRECATED = True
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
RETURN_NAMES = ("CN_WEIGHTS", "TK_SHORTCUT")
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03]
|
||||
weights = get_properly_arranged_t2i_weights(weights)
|
||||
weights = ControlWeights.t2iadapter(weights_input=weights, uncond_multiplier=uncond_multiplier, extras=cn_extras)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class CustomT2IAdapterWeightsDeprecated:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
DEPRECATED = True
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
RETURN_NAMES = ("CN_WEIGHTS", "TK_SHORTCUT")
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03]
|
||||
weights = get_properly_arranged_t2i_weights(weights)
|
||||
weights = ControlWeights.t2iadapter(weights_input=weights, uncond_multiplier=uncond_multiplier, extras=cn_extras)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class AdvancedControlNetApplyDEPR:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"control_net": ("CONTROL_NET", ),
|
||||
"image": ("IMAGE", ),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
|
||||
},
|
||||
"optional": {
|
||||
"mask_optional": ("MASK", ),
|
||||
"timestep_kf": ("TIMESTEP_KEYFRAME", ),
|
||||
"latent_kf_override": ("LATENT_KEYFRAME", ),
|
||||
"weights_override": ("CONTROL_NET_WEIGHTS", ),
|
||||
"model_optional": ("MODEL",),
|
||||
"vae_optional": ("VAE",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
DEPRECATED = True
|
||||
RETURN_TYPES = ("CONDITIONING","CONDITIONING","MODEL",)
|
||||
RETURN_NAMES = ("positive", "negative", "model_opt")
|
||||
FUNCTION = "apply_controlnet"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent,
|
||||
mask_optional=None, model_optional=None, vae_optional=None,
|
||||
timestep_kf: TimestepKeyframeGroup=None, latent_kf_override=None,
|
||||
weights_override: ControlWeights=None, control_apply_to_uncond=False):
|
||||
new_positive, new_negative = AdvancedControlNetApply.apply_controlnet(self, positive=positive, negative=negative, control_net=control_net, image=image,
|
||||
strength=strength, start_percent=start_percent, end_percent=end_percent,
|
||||
mask_optional=mask_optional, vae_optional=vae_optional,
|
||||
timestep_kf=timestep_kf, latent_kf_override=latent_kf_override, weights_override=weights_override,)
|
||||
return (new_positive, new_negative, model_optional)
|
||||
|
||||
|
||||
class AdvancedControlNetApplySingleDEPR:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"conditioning": ("CONDITIONING", ),
|
||||
"control_net": ("CONTROL_NET", ),
|
||||
"image": ("IMAGE", ),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
|
||||
},
|
||||
"optional": {
|
||||
"mask_optional": ("MASK", ),
|
||||
"timestep_kf": ("TIMESTEP_KEYFRAME", ),
|
||||
"latent_kf_override": ("LATENT_KEYFRAME", ),
|
||||
"weights_override": ("CONTROL_NET_WEIGHTS", ),
|
||||
"model_optional": ("MODEL",),
|
||||
"vae_optional": ("VAE",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
DEPRECATED = True
|
||||
RETURN_TYPES = ("CONDITIONING","MODEL",)
|
||||
RETURN_NAMES = ("CONDITIONING", "model_opt")
|
||||
FUNCTION = "apply_controlnet"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def apply_controlnet(self, conditioning, control_net, image, strength, start_percent, end_percent,
|
||||
mask_optional=None, model_optional=None, vae_optional=None,
|
||||
timestep_kf: TimestepKeyframeGroup=None, latent_kf_override=None,
|
||||
weights_override: ControlWeights=None):
|
||||
values = AdvancedControlNetApply.apply_controlnet(self, positive=conditioning, negative=None, control_net=control_net, image=image,
|
||||
strength=strength, start_percent=start_percent, end_percent=end_percent,
|
||||
mask_optional=mask_optional, vae_optional=vae_optional,
|
||||
timestep_kf=timestep_kf, latent_kf_override=latent_kf_override, weights_override=weights_override,
|
||||
control_apply_to_uncond=True)
|
||||
return (values[0], model_optional)
|
||||
|
||||
|
||||
class ControlNetLoaderAdvancedDEPR:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"control_net_name": (folder_paths.get_filename_list("controlnet"), ),
|
||||
},
|
||||
"optional": {
|
||||
"tk_optional": ("TIMESTEP_KEYFRAME", ),
|
||||
}
|
||||
}
|
||||
|
||||
DEPRECATED = True
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def load_controlnet(self, control_net_name,
|
||||
tk_optional: TimestepKeyframeGroup=None,
|
||||
timestep_keyframe: TimestepKeyframeGroup=None,
|
||||
):
|
||||
if timestep_keyframe is not None: # backwards compatibility
|
||||
tk_optional = timestep_keyframe
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
controlnet = load_controlnet(controlnet_path, tk_optional)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class DiffControlNetLoaderAdvancedDEPR:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"control_net_name": (folder_paths.get_filename_list("controlnet"), )
|
||||
},
|
||||
"optional": {
|
||||
"tk_optional": ("TIMESTEP_KEYFRAME", ),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
DEPRECATED = True
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = ""
|
||||
|
||||
def load_controlnet(self, control_net_name, model,
|
||||
tk_optional: TimestepKeyframeGroup=None,
|
||||
timestep_keyframe: TimestepKeyframeGroup=None
|
||||
):
|
||||
if timestep_keyframe is not None: # backwards compatibility
|
||||
tk_optional = timestep_keyframe
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", control_net_name)
|
||||
controlnet = load_controlnet(controlnet_path, tk_optional, model)
|
||||
if is_advanced_controlnet(controlnet):
|
||||
controlnet.verify_all_weights()
|
||||
return (controlnet,)
|
||||
|
||||
@@ -25,9 +25,6 @@ class TimestepKeyframeNode:
|
||||
"inherit_missing": ("BOOLEAN", {"default": True}, ),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
"mask_optional": ("MASK", ),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -84,9 +81,6 @@ class TimestepKeyframeInterpolationNode:
|
||||
"inherit_missing": ("BOOLEAN", {"default": True},),
|
||||
"mask_optional": ("MASK", ),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -145,9 +139,6 @@ class TimestepKeyframeFromStrengthListNode:
|
||||
"inherit_missing": ("BOOLEAN", {"default": True},),
|
||||
"mask_optional": ("MASK", ),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -204,9 +195,6 @@ class LatentKeyframeNode:
|
||||
},
|
||||
"optional": {
|
||||
"prev_latent_kf": ("LATENT_KEYFRAME", ),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -242,10 +230,7 @@ class LatentKeyframeGroupNode:
|
||||
"optional": {
|
||||
"prev_latent_kf": ("LATENT_KEYFRAME", ),
|
||||
"latent_optional": ("LATENT", ),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -364,10 +349,7 @@ class LatentKeyframeInterpolationNode:
|
||||
},
|
||||
"optional": {
|
||||
"prev_latent_kf": ("LATENT_KEYFRAME", ),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -437,10 +419,7 @@ class LatentKeyframeBatchedGroupNode:
|
||||
},
|
||||
"optional": {
|
||||
"prev_latent_kf": ("LATENT_KEYFRAME", ),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,212 +0,0 @@
|
||||
from torch import Tensor
|
||||
|
||||
import folder_paths
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .control import load_controlnet, convert_to_advanced, is_advanced_controlnet, is_sd3_advanced_controlnet
|
||||
from .utils import ControlWeights, LatentKeyframeGroup, TimestepKeyframeGroup, AbstractPreprocWrapper, BIGMAX
|
||||
|
||||
from .logger import logger
|
||||
|
||||
|
||||
class ControlNetLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"cnet": (folder_paths.get_filename_list("controlnet"), ),
|
||||
},
|
||||
"optional": {
|
||||
"_tk_opt": ("TIMESTEP_KEYFRAME", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝"
|
||||
|
||||
def load_controlnet(self, cnet,
|
||||
_tk_opt: TimestepKeyframeGroup=None,
|
||||
):
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", cnet)
|
||||
controlnet = load_controlnet(controlnet_path, _tk_opt)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class DiffControlNetLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"cnet": (folder_paths.get_filename_list("controlnet"), )
|
||||
},
|
||||
"optional": {
|
||||
"_tk_opt": ("TIMESTEP_KEYFRAME", ),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝"
|
||||
|
||||
def load_controlnet(self, cnet, model,
|
||||
_tk_opt: TimestepKeyframeGroup=None,
|
||||
):
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", cnet)
|
||||
controlnet = load_controlnet(controlnet_path, _tk_opt, model)
|
||||
if is_advanced_controlnet(controlnet):
|
||||
controlnet.verify_all_weights()
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class AdvancedControlNetApply:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"positive": ("CONDITIONING", ),
|
||||
"negative": ("CONDITIONING", ),
|
||||
"control_net": ("CONTROL_NET", ),
|
||||
"image": ("IMAGE", ),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
|
||||
},
|
||||
"optional": {
|
||||
"mask_optional": ("MASK", ),
|
||||
"timestep_kf": ("TIMESTEP_KEYFRAME", ),
|
||||
"latent_kf_override": ("LATENT_KEYFRAME", ),
|
||||
"weights_override": ("CONTROL_NET_WEIGHTS", ),
|
||||
"vae_optional": ("VAE",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING","CONDITIONING",)
|
||||
RETURN_NAMES = ("positive", "negative")
|
||||
FUNCTION = "apply_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝"
|
||||
|
||||
def apply_controlnet(self, positive, negative, control_net, image, strength, start_percent, end_percent,
|
||||
mask_optional: Tensor=None, vae_optional=None,
|
||||
timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None,
|
||||
weights_override: ControlWeights=None, control_apply_to_uncond=False):
|
||||
if strength == 0:
|
||||
return (positive, negative)
|
||||
|
||||
control_hint = image.movedim(-1,1)
|
||||
cnets = {}
|
||||
|
||||
out = []
|
||||
for conditioning in [positive, negative]:
|
||||
c = []
|
||||
if conditioning is not None:
|
||||
for t in conditioning:
|
||||
d = t[1].copy()
|
||||
|
||||
prev_cnet = d.get('control', None)
|
||||
if prev_cnet in cnets:
|
||||
c_net = cnets[prev_cnet]
|
||||
else:
|
||||
# make sure control_net is not None to avoid confusing error messages
|
||||
if control_net is None:
|
||||
raise Exception("Passed in control_net is None; something must have went wrong when loading it from a Load ControlNet node.")
|
||||
# copy, convert to advanced if needed, and set cond
|
||||
c_net = convert_to_advanced(control_net.copy()).set_cond_hint(control_hint, strength, (start_percent, end_percent), vae_optional)
|
||||
if is_advanced_controlnet(c_net):
|
||||
# disarm node check
|
||||
c_net.disarm()
|
||||
# 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
|
||||
if c_net.allow_condhint_latents and isinstance(control_hint, AbstractPreprocWrapper):
|
||||
pass
|
||||
elif not vae_optional:
|
||||
# make sure SD3 ControlNet will get a special message instead of generic type mention
|
||||
if is_sd3_advanced_controlnet(c_net):
|
||||
raise Exception(f"SD3 ControlNet requires vae_optional input, but got None.")
|
||||
else:
|
||||
raise Exception(f"Type '{type(c_net).__name__}' requires vae_optional input, but got None.")
|
||||
# apply optional parameters and overrides, if provided
|
||||
if timestep_kf is not None:
|
||||
c_net.set_timestep_keyframes(timestep_kf)
|
||||
if latent_kf_override is not None:
|
||||
c_net.latent_keyframe_override = latent_kf_override
|
||||
if weights_override is not None:
|
||||
c_net.weights_override = weights_override
|
||||
# verify weights are compatible
|
||||
c_net.verify_all_weights()
|
||||
# set cond hint mask
|
||||
if mask_optional is not None:
|
||||
mask_optional = mask_optional.clone()
|
||||
# if not in the form of a batch, make it so
|
||||
if len(mask_optional.shape) < 3:
|
||||
mask_optional = mask_optional.unsqueeze(0)
|
||||
c_net.set_cond_hint_mask(mask_optional)
|
||||
c_net.set_previous_controlnet(prev_cnet)
|
||||
cnets[prev_cnet] = c_net
|
||||
|
||||
d['control'] = c_net
|
||||
d['control_apply_to_uncond'] = control_apply_to_uncond
|
||||
n = [t[0], d]
|
||||
c.append(n)
|
||||
out.append(c)
|
||||
return (out[0], out[1])
|
||||
|
||||
|
||||
class AdvancedControlNetApplySingle:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"conditioning": ("CONDITIONING", ),
|
||||
"control_net": ("CONTROL_NET", ),
|
||||
"image": ("IMAGE", ),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
|
||||
},
|
||||
"optional": {
|
||||
"mask_optional": ("MASK", ),
|
||||
"timestep_kf": ("TIMESTEP_KEYFRAME", ),
|
||||
"latent_kf_override": ("LATENT_KEYFRAME", ),
|
||||
"weights_override": ("CONTROL_NET_WEIGHTS", ),
|
||||
"vae_optional": ("VAE",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING","MODEL",)
|
||||
RETURN_NAMES = ("CONDITIONING", "model_opt")
|
||||
FUNCTION = "apply_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝"
|
||||
|
||||
def apply_controlnet(self, conditioning, control_net, image, strength, start_percent, end_percent,
|
||||
mask_optional: Tensor=None, vae_optional=None,
|
||||
timestep_kf: TimestepKeyframeGroup=None, latent_kf_override: LatentKeyframeGroup=None,
|
||||
weights_override: ControlWeights=None):
|
||||
values = AdvancedControlNetApply.apply_controlnet(self, positive=conditioning, negative=None, control_net=control_net, image=image,
|
||||
strength=strength, start_percent=start_percent, end_percent=end_percent,
|
||||
mask_optional=mask_optional, vae_optional=vae_optional,
|
||||
timestep_kf=timestep_kf, latent_kf_override=latent_kf_override, weights_override=weights_override,
|
||||
control_apply_to_uncond=True)
|
||||
return (values[0],)
|
||||
@@ -1,88 +0,0 @@
|
||||
from torch import Tensor
|
||||
import math
|
||||
|
||||
import folder_paths
|
||||
|
||||
from .control_plusplus import load_controlnetplusplus, PlusPlusType, PlusPlusInput, PlusPlusInputGroup, PlusPlusImageWrapper
|
||||
from .utils import BIGMAX
|
||||
|
||||
|
||||
class PlusPlusLoaderAdvanced:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"plus_input": ("PLUS_INPUT", ),
|
||||
"name": (folder_paths.get_filename_list("controlnet"), ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", "IMAGE",)
|
||||
FUNCTION = "load_controlnet_plusplus"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/ControlNet++"
|
||||
|
||||
def load_controlnet_plusplus(self, plus_input: PlusPlusInputGroup, name: str):
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", name)
|
||||
controlnet = load_controlnetplusplus(controlnet_path)
|
||||
controlnet.verify_control_type(name, plus_input)
|
||||
controlnet.allow_condhint_latents = True
|
||||
return (controlnet, PlusPlusImageWrapper(plus_input),)
|
||||
|
||||
|
||||
class PlusPlusLoaderSingle:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"name": (folder_paths.get_filename_list("controlnet"), ),
|
||||
"control_type": (PlusPlusType._LIST_WITH_NONE, {"default": PlusPlusType.NONE}, ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET",)
|
||||
FUNCTION = "load_controlnet_plusplus"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/ControlNet++"
|
||||
|
||||
def load_controlnet_plusplus(self, name: str, control_type: str):
|
||||
controlnet_path = folder_paths.get_full_path("controlnet", name)
|
||||
controlnet = load_controlnetplusplus(controlnet_path)
|
||||
controlnet.single_control_type = control_type
|
||||
controlnet.verify_control_type(name)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class PlusPlusInputNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",),
|
||||
"control_type": (PlusPlusType._LIST,),
|
||||
},
|
||||
"optional": {
|
||||
"prev_plus_input": ("PLUS_INPUT",),
|
||||
#"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": BIGMAX, "step": 0.01}),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PLUS_INPUT", )
|
||||
FUNCTION = "wrap_images"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/ControlNet++"
|
||||
|
||||
def wrap_images(self, image: Tensor, control_type: str, strength=1.0, prev_plus_input: PlusPlusInputGroup=None):
|
||||
if prev_plus_input is None:
|
||||
prev_plus_input = PlusPlusInputGroup()
|
||||
prev_plus_input = prev_plus_input.clone()
|
||||
|
||||
if math.isclose(strength, 0.0):
|
||||
strength = 0.0000001
|
||||
pp_input = PlusPlusInput(image, control_type, strength)
|
||||
prev_plus_input.add(pp_input)
|
||||
|
||||
return (prev_plus_input,)
|
||||
@@ -6,7 +6,7 @@ import comfy.utils
|
||||
from comfy.sd import VAE
|
||||
|
||||
from .utils import TimestepKeyframeGroup
|
||||
from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper, SparseConst, SparseContextAware, get_idx_list_from_str
|
||||
from .control_sparsectrl import SparseMethod, SparseIndexMethod, SparseSettings, SparseSpreadMethod, PreprocSparseRGBWrapper
|
||||
from .control import load_sparsectrl, load_controlnet, ControlNetAdvanced, SparseCtrlAdvanced
|
||||
|
||||
|
||||
@@ -24,10 +24,6 @@ class SparseCtrlLoaderAdvanced:
|
||||
"optional": {
|
||||
"sparse_method": ("SPARSE_METHOD", ),
|
||||
"tk_optional": ("TIMESTEP_KEYFRAME", ),
|
||||
"context_aware": (SparseContextAware.LIST, ),
|
||||
"sparse_hint_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_nonhint_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_mask_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -36,12 +32,9 @@ class SparseCtrlLoaderAdvanced:
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl"
|
||||
|
||||
def load_controlnet(self, sparsectrl_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None,
|
||||
context_aware=SparseContextAware.NEAREST_HINT, sparse_hint_mult=1.0, sparse_nonhint_mult=1.0, sparse_mask_mult=1.0):
|
||||
def load_controlnet(self, sparsectrl_name: str, use_motion: bool, motion_strength: float, motion_scale: float, sparse_method: SparseMethod=SparseSpreadMethod(), tk_optional: TimestepKeyframeGroup=None):
|
||||
sparsectrl_path = folder_paths.get_full_path("controlnet", sparsectrl_name)
|
||||
sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale,
|
||||
context_aware=context_aware,
|
||||
sparse_mask_mult=sparse_mask_mult, sparse_hint_mult=sparse_hint_mult, sparse_nonhint_mult=sparse_nonhint_mult)
|
||||
sparse_settings = SparseSettings(sparse_method=sparse_method, use_motion=use_motion, motion_strength=motion_strength, motion_scale=motion_scale)
|
||||
sparsectrl = load_sparsectrl(sparsectrl_path, timestep_keyframe=tk_optional, sparse_settings=sparse_settings)
|
||||
return (sparsectrl,)
|
||||
|
||||
@@ -103,7 +96,21 @@ class SparseIndexMethodNode:
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl"
|
||||
|
||||
def get_method(self, indexes: str):
|
||||
idxs = get_idx_list_from_str(indexes)
|
||||
idxs = []
|
||||
unique_idxs = set()
|
||||
# get indeces from string
|
||||
str_idxs = [x.strip() for x in indexes.strip().split(",")]
|
||||
for str_idx in str_idxs:
|
||||
try:
|
||||
idx = int(str_idx)
|
||||
if idx in unique_idxs:
|
||||
raise ValueError(f"'{idx}' is duplicated; indexes must be unique.")
|
||||
idxs.append(idx)
|
||||
unique_idxs.add(idx)
|
||||
except ValueError:
|
||||
raise ValueError(f"'{str_idx}' is not a valid integer index.")
|
||||
if len(idxs) == 0:
|
||||
raise ValueError(f"No indexes were listed in Sparse Index Method.")
|
||||
return (SparseIndexMethod(idxs),)
|
||||
|
||||
|
||||
@@ -133,9 +140,6 @@ class RgbSparseCtrlPreprocessor:
|
||||
"image": ("IMAGE", ),
|
||||
"vae": ("VAE", ),
|
||||
"latent_size": ("LATENT", ),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -157,32 +161,3 @@ class RgbSparseCtrlPreprocessor:
|
||||
image = VAEEncode.vae_encode_crop_pixels(image)
|
||||
encoded = vae.encode(image[:,:,:,:3])
|
||||
return (PreprocSparseRGBWrapper(condhint=encoded),)
|
||||
|
||||
|
||||
class SparseWeightExtras:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"optional": {
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
"sparse_hint_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_nonhint_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"sparse_mask_mult": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CN_WEIGHTS_EXTRAS", )
|
||||
RETURN_NAMES = ("cn_extras", )
|
||||
FUNCTION = "create_weight_extras"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/SparseCtrl/extras"
|
||||
|
||||
def create_weight_extras(self, cn_extras: dict[str]={}, sparse_hint_mult=1.0, sparse_nonhint_mult=1.0, sparse_mask_mult=1.0):
|
||||
cn_extras = cn_extras.copy()
|
||||
cn_extras[SparseConst.HINT_MULT] = sparse_hint_mult
|
||||
cn_extras[SparseConst.NONHINT_MULT] = sparse_nonhint_mult
|
||||
cn_extras[SparseConst.MASK_MULT] = sparse_mask_mult
|
||||
return (cn_extras, )
|
||||
|
||||
+89
-194
@@ -1,6 +1,6 @@
|
||||
from torch import Tensor
|
||||
import torch
|
||||
from .utils import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, Extras, get_properly_arranged_t2i_weights, linear_conversion
|
||||
from .utils import TimestepKeyframe, TimestepKeyframeGroup, ControlWeights, get_properly_arranged_t2i_weights, linear_conversion
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -11,12 +11,6 @@ class DefaultWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"optional": {
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
@@ -25,8 +19,8 @@ class DefaultWeights:
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights"
|
||||
|
||||
def load_weights(self, cn_extras: dict[str]={}):
|
||||
weights = ControlWeights.default(extras=cn_extras)
|
||||
def load_weights(self):
|
||||
weights = ControlWeights.default()
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
@@ -43,10 +37,6 @@ class ScaledSoftMaskedUniversalWeights:
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -57,7 +47,7 @@ class ScaledSoftMaskedUniversalWeights:
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights"
|
||||
|
||||
def load_weights(self, mask: Tensor, min_base_multiplier: float, max_base_multiplier: float, lock_min=False, lock_max=False,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
uncond_multiplier: float=1.0):
|
||||
# normalize mask
|
||||
mask = mask.clone()
|
||||
x_min = 0.0 if lock_min else mask.min()
|
||||
@@ -66,7 +56,7 @@ class ScaledSoftMaskedUniversalWeights:
|
||||
mask = torch.ones_like(mask) * max_base_multiplier
|
||||
else:
|
||||
mask = linear_conversion(mask, x_min, x_max, min_base_multiplier, max_base_multiplier)
|
||||
weights = ControlWeights.universal_mask(weight_mask=mask, uncond_multiplier=uncond_multiplier, extras=cn_extras)
|
||||
weights = ControlWeights.universal_mask(weight_mask=mask, uncond_multiplier=uncond_multiplier)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
@@ -76,13 +66,10 @@ class ScaledSoftUniversalWeights:
|
||||
return {
|
||||
"required": {
|
||||
"base_multiplier": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 1.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -92,36 +79,73 @@ class ScaledSoftUniversalWeights:
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights"
|
||||
|
||||
def load_weights(self, base_multiplier, uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights = ControlWeights.universal(base_multiplier=base_multiplier, uncond_multiplier=uncond_multiplier, extras=cn_extras)
|
||||
def load_weights(self, base_multiplier, flip_weights, uncond_multiplier: float=1.0):
|
||||
weights = ControlWeights.universal(base_multiplier=base_multiplier, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class SoftControlNetWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"weight_00": ("FLOAT", {"default": 0.09941396206337118, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 0.12050177219802567, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 0.14606275417942507, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 0.17704576264172736, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_04": ("FLOAT", {"default": 0.214600924414215, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_05": ("FLOAT", {"default": 0.26012233262329093, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_06": ("FLOAT", {"default": 0.3152997971191405, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_07": ("FLOAT", {"default": 0.3821815722656249, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_08": ("FLOAT", {"default": 0.4632503906249999, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_09": ("FLOAT", {"default": 0.561515625, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_10": ("FLOAT", {"default": 0.6806249999999999, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_11": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
RETURN_NAMES = WEIGHTS_RETURN_NAMES
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet"
|
||||
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights,
|
||||
uncond_multiplier: float=1.0):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12]
|
||||
weights = ControlWeights.controlnet(weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class SoftControlNetWeightsSD15:
|
||||
class CustomControlNetWeights:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"output_0": ("FLOAT", {"default": 0.09941396206337118, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_1": ("FLOAT", {"default": 0.12050177219802567, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_2": ("FLOAT", {"default": 0.14606275417942507, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_3": ("FLOAT", {"default": 0.17704576264172736, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_4": ("FLOAT", {"default": 0.214600924414215, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_5": ("FLOAT", {"default": 0.26012233262329093, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_6": ("FLOAT", {"default": 0.3152997971191405, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_7": ("FLOAT", {"default": 0.3821815722656249, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_8": ("FLOAT", {"default": 0.4632503906249999, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_9": ("FLOAT", {"default": 0.561515625, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_10": ("FLOAT", {"default": 0.6806249999999999, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_11": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"middle_0": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_04": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_05": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_06": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_07": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_08": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_09": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_10": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_11": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -131,110 +155,12 @@ class SoftControlNetWeightsSD15:
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet"
|
||||
|
||||
def load_weights(self, output_0, output_1, output_2, output_3, output_4, output_5, output_6,
|
||||
output_7, output_8, output_9, output_10, output_11, middle_0,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
return CustomControlNetWeightsSD15.load_weights(self,
|
||||
output_0=output_0, output_1=output_1, output_2=output_2, output_3=output_3,
|
||||
output_4=output_4, output_5=output_5, output_6=output_6, output_7=output_7,
|
||||
output_8=output_8, output_9=output_9, output_10=output_10, output_11=output_11,
|
||||
middle_0=middle_0,
|
||||
uncond_multiplier=uncond_multiplier, cn_extras=cn_extras)
|
||||
|
||||
|
||||
class CustomControlNetWeightsSD15:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"output_0": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_1": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_2": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_3": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_4": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_5": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_6": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_7": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_8": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_9": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_10": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"output_11": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"middle_0": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
RETURN_NAMES = WEIGHTS_RETURN_NAMES
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet"
|
||||
|
||||
def load_weights(self, output_0, output_1, output_2, output_3, output_4, output_5, output_6,
|
||||
output_7, output_8, output_9, output_10, output_11, middle_0,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights_output = [output_0, output_1, output_2, output_3, output_4, output_5, output_6,
|
||||
output_7, output_8, output_9, output_10, output_11]
|
||||
weights_middle = [middle_0]
|
||||
weights = ControlWeights.controlnet(weights_output=weights_output, weights_middle=weights_middle, uncond_multiplier=uncond_multiplier,
|
||||
extras=cn_extras, disable_applied_to=True)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class CustomControlNetWeightsFlux:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_0": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_1": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_2": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_3": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_4": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_5": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_6": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_7": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_8": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_9": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_10": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_11": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_12": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_13": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_14": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_15": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_16": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_17": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_18": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET_WEIGHTS", "TIMESTEP_KEYFRAME",)
|
||||
RETURN_NAMES = WEIGHTS_RETURN_NAMES
|
||||
FUNCTION = "load_weights"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/ControlNet"
|
||||
|
||||
def load_weights(self, input_0, input_1, input_2, input_3, input_4, input_5, input_6,
|
||||
input_7, input_8, input_9, input_10, input_11, input_12, input_13,
|
||||
input_14, input_15, input_16, input_17, input_18,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights_input = [input_0, input_1, input_2, input_3, input_4, input_5,
|
||||
input_6, input_7, input_8, input_9, input_10, input_11,
|
||||
input_12, input_13, input_14, input_15, input_16, input_17, input_18]
|
||||
weights = ControlWeights.controlnet(weights_input=weights_input, uncond_multiplier=uncond_multiplier, extras=cn_extras, disable_applied_to=True)
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12, flip_weights,
|
||||
uncond_multiplier: float=1.0):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03, weight_04, weight_05, weight_06,
|
||||
weight_07, weight_08, weight_09, weight_10, weight_11, weight_12]
|
||||
weights = ControlWeights.controlnet(weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
@@ -243,17 +169,14 @@ class SoftT2IAdapterWeights:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_0": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_1": ("FLOAT", {"default": 0.62, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_2": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_3": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_00": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 0.62, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 0.825, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -263,10 +186,12 @@ class SoftT2IAdapterWeights:
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter"
|
||||
|
||||
def load_weights(self, input_0, input_1, input_2, input_3,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
return CustomT2IAdapterWeights.load_weights(self, input_0=input_0, input_1=input_1, input_2=input_2, input_3=input_3,
|
||||
uncond_multiplier=uncond_multiplier, cn_extras=cn_extras)
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights,
|
||||
uncond_multiplier: float=1.0):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03]
|
||||
weights = get_properly_arranged_t2i_weights(weights)
|
||||
weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
return (weights, TimestepKeyframeGroup.default(TimestepKeyframe(control_weights=weights)))
|
||||
|
||||
|
||||
class CustomT2IAdapterWeights:
|
||||
@@ -274,17 +199,14 @@ class CustomT2IAdapterWeights:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"input_0": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_1": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_2": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"input_3": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_00": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_01": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_02": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"weight_03": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"flip_weights": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"uncond_multiplier": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}, ),
|
||||
"cn_extras": ("CN_WEIGHTS_EXTRAS",),
|
||||
},
|
||||
"hidden": {
|
||||
"autosize": ("ACNAUTOSIZE", {"padding": 0}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -294,36 +216,9 @@ class CustomT2IAdapterWeights:
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/weights/T2IAdapter"
|
||||
|
||||
def load_weights(self, input_0, input_1, input_2, input_3,
|
||||
uncond_multiplier: float=1.0, cn_extras: dict[str]={}):
|
||||
weights = [input_0, input_1, input_2, input_3]
|
||||
def load_weights(self, weight_00, weight_01, weight_02, weight_03, flip_weights,
|
||||
uncond_multiplier: float=1.0):
|
||||
weights = [weight_00, weight_01, weight_02, weight_03]
|
||||
weights = get_properly_arranged_t2i_weights(weights)
|
||||
weights = ControlWeights.t2iadapter(weights_input=weights, uncond_multiplier=uncond_multiplier, extras=cn_extras, disable_applied_to=True)
|
||||
weights = ControlWeights.t2iadapter(weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
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,)
|
||||
|
||||
@@ -1,225 +0,0 @@
|
||||
from typing import Callable, Union
|
||||
|
||||
import comfy.hooks
|
||||
import comfy.model_patcher
|
||||
import comfy.patcher_extension
|
||||
import comfy.sample
|
||||
import comfy.samplers
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.controlnet import ControlBase
|
||||
from comfy.ldm.modules.attention import BasicTransformerBlock
|
||||
|
||||
|
||||
from .control import convert_all_to_advanced, restore_all_controlnet_conns
|
||||
from .control_reference import (ReferenceAdvanced, ReferenceInjections,
|
||||
RefBasicTransformerBlock, RefTimestepEmbedSequential,
|
||||
InjectionBasicTransformerBlockHolder, InjectionTimestepEmbedSequentialHolder,
|
||||
_forward_inject_BasicTransformerBlock,
|
||||
handle_context_ref_setup, handle_reference_injection,
|
||||
REF_CONTROL_LIST_ALL, CONTEXTREF_CLEAN_FUNC)
|
||||
from .dinklink import get_dinklink
|
||||
from .utils import torch_dfs, WrapperConsts, CURRENT_WRAPPER_VERSION
|
||||
|
||||
def prepare_dinklink_acn_wrapper():
|
||||
# expose acn_sampler_sample_wrapper
|
||||
d = get_dinklink()
|
||||
link_acn = d.setdefault(WrapperConsts.ACN, {})
|
||||
link_acn[WrapperConsts.VERSION] = CURRENT_WRAPPER_VERSION
|
||||
link_acn[WrapperConsts.ACN_CREATE_SAMPLER_SAMPLE_WRAPPER] = (comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
|
||||
WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY,
|
||||
acn_outer_sample_wrapper)
|
||||
|
||||
|
||||
def support_sliding_context_windows(conds) -> tuple[bool, list[dict]]:
|
||||
# convert to advanced, with report if anything was actually modified
|
||||
modified, new_conds = convert_all_to_advanced(conds)
|
||||
return modified, new_conds
|
||||
|
||||
|
||||
def has_sliding_context_windows(model: ModelPatcher):
|
||||
params = model.get_attachment("ADE_params")
|
||||
if params is None:
|
||||
# backwards compatibility
|
||||
params = getattr(model, "motion_injection_params", None)
|
||||
if params is None:
|
||||
return False
|
||||
context_options = getattr(params, "context_options")
|
||||
return context_options.context_length is not None
|
||||
|
||||
|
||||
def get_contextref_obj(model: ModelPatcher):
|
||||
params = model.get_attachment("ADE_params")
|
||||
if params is None:
|
||||
# backwards compatibility
|
||||
params = getattr(model, "motion_injection_params", None)
|
||||
if params is None:
|
||||
return None
|
||||
context_options = getattr(params, "context_options")
|
||||
extras = getattr(context_options, "extras", None)
|
||||
if extras is None:
|
||||
return None
|
||||
return getattr(extras, "context_ref", None)
|
||||
|
||||
|
||||
def get_refcn(control: ControlBase, order: int=-1):
|
||||
ref_set: set[ReferenceAdvanced] = set()
|
||||
if control is None:
|
||||
return ref_set
|
||||
if type(control) == ReferenceAdvanced and not control.is_context_ref:
|
||||
control.order = order
|
||||
order -= 1
|
||||
ref_set.add(control)
|
||||
ref_set.update(get_refcn(control.previous_controlnet, order=order))
|
||||
return ref_set
|
||||
|
||||
|
||||
def should_register_outer_sample_wrapper(hook, model, model_options: dict, target, registered: list):
|
||||
wrappers = comfy.patcher_extension.get_wrappers_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
|
||||
WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY,
|
||||
model_options, is_model_options=True)
|
||||
return len(wrappers) == 0
|
||||
|
||||
def create_wrapper_hooks():
|
||||
wrappers = {}
|
||||
comfy.patcher_extension.add_wrapper_with_key(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE,
|
||||
WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY,
|
||||
acn_outer_sample_wrapper,
|
||||
transformer_options=wrappers)
|
||||
hooks = comfy.hooks.HookGroup()
|
||||
hook = comfy.hooks.WrapperHook(wrappers)
|
||||
hook.hook_id = WrapperConsts.ACN_OUTER_SAMPLE_WRAPPER_KEY
|
||||
hook.custom_should_register = should_register_outer_sample_wrapper
|
||||
hooks.add(hook)
|
||||
return hooks
|
||||
|
||||
def acn_outer_sample_wrapper(executor, *args, **kwargs):
|
||||
controlnets_modified = False
|
||||
guider: comfy.samplers.CFGGuider = executor.class_obj
|
||||
model = guider.model_patcher
|
||||
orig_conds = guider.conds
|
||||
orig_model_options = guider.model_options
|
||||
try:
|
||||
new_model_options = orig_model_options
|
||||
# if context options present, perform some special actions that may be required
|
||||
context_refs = []
|
||||
if has_sliding_context_windows(guider.model_patcher):
|
||||
new_model_options = comfy.model_patcher.create_model_options_clone(new_model_options)
|
||||
# convert all CNs to Advanced if needed
|
||||
controlnets_modified, conds = support_sliding_context_windows(orig_conds)
|
||||
if controlnets_modified:
|
||||
guider.conds = conds
|
||||
# enable ContextRef, if requested
|
||||
existing_contextref_obj = get_contextref_obj(guider.model_patcher)
|
||||
if existing_contextref_obj is not None:
|
||||
context_refs = handle_context_ref_setup(existing_contextref_obj, new_model_options["transformer_options"], guider.conds)
|
||||
controlnets_modified = True
|
||||
# look for Advanced ControlNets that will require intervention to work
|
||||
ref_set = set()
|
||||
for outer_cond in guider.conds.values():
|
||||
for cond in outer_cond:
|
||||
if "control" in cond:
|
||||
ref_set.update(get_refcn(cond["control"]))
|
||||
# if no ref cn found, do original function immediately
|
||||
if len(ref_set) == 0 and len(context_refs) == 0:
|
||||
return executor(*args, **kwargs)
|
||||
# otherwise, injection time
|
||||
try:
|
||||
# inject
|
||||
# storage for all Reference-related injections
|
||||
reference_injections = ReferenceInjections()
|
||||
|
||||
# first, handle attn module injection
|
||||
all_modules = torch_dfs(model.model)
|
||||
attn_modules: list[RefBasicTransformerBlock] = []
|
||||
for module in all_modules:
|
||||
if isinstance(module, BasicTransformerBlock):
|
||||
attn_modules.append(module)
|
||||
attn_modules = [module for module in all_modules if isinstance(module, BasicTransformerBlock)]
|
||||
attn_modules = sorted(attn_modules, key=lambda x: -x.norm1.normalized_shape[0])
|
||||
for i, module in enumerate(attn_modules):
|
||||
injection_holder = InjectionBasicTransformerBlockHolder(block=module, idx=i)
|
||||
injection_holder.attn_weight = float(i) / float(len(attn_modules))
|
||||
if hasattr(module, "_forward"): # backward compatibility
|
||||
module._forward = _forward_inject_BasicTransformerBlock.__get__(module, type(module))
|
||||
else:
|
||||
module.forward = _forward_inject_BasicTransformerBlock.__get__(module, type(module))
|
||||
module.injection_holder = injection_holder
|
||||
reference_injections.attn_modules.append(module)
|
||||
# figure out which module is middle block
|
||||
if hasattr(model.model.diffusion_model, "middle_block"):
|
||||
mid_modules = torch_dfs(model.model.diffusion_model.middle_block)
|
||||
mid_attn_modules: list[RefBasicTransformerBlock] = [module for module in mid_modules if isinstance(module, BasicTransformerBlock)]
|
||||
for module in mid_attn_modules:
|
||||
module.injection_holder.is_middle = True
|
||||
|
||||
# next, handle gn module injection (TimestepEmbedSequential)
|
||||
# TODO: figure out the logic behind these hardcoded indexes
|
||||
if type(model.model).__name__ == "SDXL":
|
||||
input_block_indices = [4, 5, 7, 8]
|
||||
output_block_indices = [0, 1, 2, 3, 4, 5]
|
||||
else:
|
||||
input_block_indices = [4, 5, 7, 8, 10, 11]
|
||||
output_block_indices = [0, 1, 2, 3, 4, 5, 6, 7]
|
||||
if hasattr(model.model.diffusion_model, "middle_block"):
|
||||
module = model.model.diffusion_model.middle_block
|
||||
injection_holder = InjectionTimestepEmbedSequentialHolder(block=module, idx=0, is_middle=True)
|
||||
injection_holder.gn_weight = 0.0
|
||||
module.injection_holder = injection_holder
|
||||
reference_injections.gn_modules.append(module)
|
||||
for w, i in enumerate(input_block_indices):
|
||||
module = model.model.diffusion_model.input_blocks[i]
|
||||
injection_holder = InjectionTimestepEmbedSequentialHolder(block=module, idx=i, is_input=True)
|
||||
injection_holder.gn_weight = 1.0 - float(w) / float(len(input_block_indices))
|
||||
module.injection_holder = injection_holder
|
||||
reference_injections.gn_modules.append(module)
|
||||
for w, i in enumerate(output_block_indices):
|
||||
module = model.model.diffusion_model.output_blocks[i]
|
||||
injection_holder = InjectionTimestepEmbedSequentialHolder(block=module, idx=i, is_output=True)
|
||||
injection_holder.gn_weight = float(w) / float(len(output_block_indices))
|
||||
module.injection_holder = injection_holder
|
||||
reference_injections.gn_modules.append(module)
|
||||
# hack gn_module forwards and update weights
|
||||
for i, module in enumerate(reference_injections.gn_modules):
|
||||
module.injection_holder.gn_weight *= 2
|
||||
|
||||
# store ordered ref cns in model's transformer options
|
||||
new_model_options = comfy.model_patcher.create_model_options_clone(new_model_options)
|
||||
# handle diffusion_model forward injection
|
||||
handle_reference_injection(new_model_options, reference_injections)
|
||||
|
||||
ref_list: list[ReferenceAdvanced] = list(ref_set)
|
||||
new_model_options["transformer_options"][REF_CONTROL_LIST_ALL] = sorted(ref_list, key=lambda x: x.order)
|
||||
new_model_options["transformer_options"][CONTEXTREF_CLEAN_FUNC] = reference_injections.clean_contextref_module_mem
|
||||
guider.model_options = new_model_options
|
||||
# continue with original function
|
||||
return executor(*args, **kwargs)
|
||||
finally:
|
||||
# cleanup injections
|
||||
# restore attn modules
|
||||
attn_modules: list[RefBasicTransformerBlock] = reference_injections.attn_modules
|
||||
for module in attn_modules:
|
||||
module.injection_holder.restore(module)
|
||||
module.injection_holder.clean_all()
|
||||
del module.injection_holder
|
||||
del attn_modules
|
||||
# restore gn modules
|
||||
gn_modules: list[RefTimestepEmbedSequential] = reference_injections.gn_modules
|
||||
for module in gn_modules:
|
||||
module.injection_holder.restore(module)
|
||||
module.injection_holder.clean_all()
|
||||
del module.injection_holder
|
||||
del gn_modules
|
||||
# cleanup
|
||||
reference_injections.cleanup()
|
||||
finally:
|
||||
# restore model_options
|
||||
guider.model_options = orig_model_options
|
||||
# restore guider.conds
|
||||
guider.conds = orig_conds
|
||||
# restore controlnets in conds, if needed
|
||||
if controlnets_modified:
|
||||
restore_all_controlnet_conns(guider.conds)
|
||||
del orig_conds
|
||||
del orig_model_options
|
||||
del model
|
||||
del guider
|
||||
+212
-212
@@ -3,29 +3,22 @@ from typing import Callable, Union
|
||||
import torch
|
||||
from torch import Tensor
|
||||
import torch.nn.functional
|
||||
from einops import rearrange
|
||||
import numpy as np
|
||||
import math
|
||||
|
||||
import comfy.ops
|
||||
import comfy.utils
|
||||
import comfy.sample
|
||||
import comfy.samplers
|
||||
import comfy.model_base
|
||||
|
||||
from comfy.controlnet import ControlBase
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.sd import VAE
|
||||
|
||||
from .logger import logger
|
||||
|
||||
BIGMIN = -(2**53-1)
|
||||
BIGMAX = (2**53-1)
|
||||
BIGMAX_TENSOR = torch.tensor(9999999999.9)
|
||||
|
||||
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):
|
||||
@@ -34,14 +27,107 @@ def load_torch_file_with_dict_factory(controlnet_data: dict[str, Tensor], orig_l
|
||||
return controlnet_data
|
||||
return load_torch_file_with_dict
|
||||
|
||||
# wrapping len function so that it will save the thing len is trying to get the length of;
|
||||
# this will be assumed to be the cond_or_uncond variable;
|
||||
# automatically restores len to original function after running
|
||||
def wrapper_len_factory(orig_len: Callable) -> Callable:
|
||||
def wrapper_len(*args, **kwargs):
|
||||
cond_or_uncond = args[0]
|
||||
real_length = orig_len(*args, **kwargs)
|
||||
if real_length > 0 and type(cond_or_uncond) == list and (cond_or_uncond[0] in [0, 1]):
|
||||
try:
|
||||
to_return = IntWithCondOrUncond(real_length)
|
||||
setattr(to_return, "cond_or_uncond", cond_or_uncond)
|
||||
return to_return
|
||||
finally:
|
||||
__builtins__["len"] = orig_len
|
||||
else:
|
||||
return real_length
|
||||
return wrapper_len
|
||||
|
||||
CURRENT_WRAPPER_VERSION = 10002
|
||||
# wrapping cond_cat function so that it will wrap around len function to get cond_or_uncond variable value
|
||||
# from comfy.samplers.calc_conds_batch
|
||||
def wrapper_cond_cat_factory(orig_cond_cat: Callable):
|
||||
def wrapper_cond_cat(*args, **kwargs):
|
||||
__builtins__["len"] = wrapper_len_factory(__builtins__["len"])
|
||||
return orig_cond_cat(*args, **kwargs)
|
||||
return wrapper_cond_cat
|
||||
orig_cond_cat = comfy.samplers.cond_cat
|
||||
comfy.samplers.cond_cat = wrapper_cond_cat_factory(orig_cond_cat)
|
||||
|
||||
|
||||
# wrapping apply_model so that len function will be cleaned up fairly soon after being injected
|
||||
def apply_model_uncond_cleanup_factory(orig_apply_model, orig_len):
|
||||
def apply_model_uncond_cleanup_wrapper(self, *args, **kwargs):
|
||||
__builtins__["len"] = orig_len
|
||||
return orig_apply_model(self, *args, **kwargs)
|
||||
return apply_model_uncond_cleanup_wrapper
|
||||
global_orig_len = __builtins__["len"]
|
||||
orig_apply_model = comfy.model_base.BaseModel.apply_model
|
||||
comfy.model_base.BaseModel.apply_model = apply_model_uncond_cleanup_factory(orig_apply_model, global_orig_len)
|
||||
|
||||
|
||||
def uncond_multiplier_check_cn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable:
|
||||
def contains_uncond_multiplier(control: Union[ControlBase, 'AdvancedControlBase']):
|
||||
if control is None:
|
||||
return False
|
||||
if not isinstance(control, AdvancedControlBase):
|
||||
return contains_uncond_multiplier(control.previous_controlnet)
|
||||
# check if weights_override has an uncond_multiplier
|
||||
if control.weights_override is not None and control.weights_override.has_uncond_multiplier:
|
||||
return True
|
||||
# check if any timestep_keyframes have an uncond_multiplier on their weights
|
||||
if control.timestep_keyframes is not None:
|
||||
for tk in control.timestep_keyframes.keyframes:
|
||||
if tk.has_control_weights() and tk.control_weights.has_uncond_multiplier:
|
||||
return True
|
||||
return contains_uncond_multiplier(control.previous_controlnet)
|
||||
|
||||
# check if positive or negative conds contain Adv. Cns that use multiply_negative on weights
|
||||
def uncond_multiplier_check_cn_sample(model: ModelPatcher, *args, **kwargs):
|
||||
positive = args[-3]
|
||||
negative = args[-2]
|
||||
has_uncond_multiplier = False
|
||||
if positive is not None:
|
||||
for cond in positive:
|
||||
if "control" in cond[1]:
|
||||
has_uncond_multiplier = contains_uncond_multiplier(cond[1]["control"])
|
||||
if has_uncond_multiplier:
|
||||
break
|
||||
if negative is not None and not has_uncond_multiplier:
|
||||
for cond in negative:
|
||||
if "control" in cond[1]:
|
||||
has_uncond_multiplier = contains_uncond_multiplier(cond[1]["control"])
|
||||
if has_uncond_multiplier:
|
||||
break
|
||||
try:
|
||||
# if uncond_multiplier found, continue to use wrapped version of function
|
||||
if has_uncond_multiplier:
|
||||
return orig_comfy_sample(model, *args, **kwargs)
|
||||
# otherwise, use original version of function to prevent even the smallest of slowdowns (0.XX%)
|
||||
try:
|
||||
wrapped_cond_cat = comfy.samplers.cond_cat
|
||||
comfy.samplers.cond_cat = orig_cond_cat
|
||||
return orig_comfy_sample(model, *args, **kwargs)
|
||||
finally:
|
||||
comfy.samplers.cond_cat = wrapped_cond_cat
|
||||
finally:
|
||||
# make sure len function is unwrapped by the time sampling is done, just in case
|
||||
__builtins__["len"] = global_orig_len
|
||||
return uncond_multiplier_check_cn_sample
|
||||
# inject sample functions
|
||||
comfy.sample.sample = uncond_multiplier_check_cn_sample_factory(comfy.sample.sample)
|
||||
comfy.sample.sample_custom = uncond_multiplier_check_cn_sample_factory(comfy.sample.sample_custom, is_custom=True)
|
||||
|
||||
|
||||
class IntWithCondOrUncond(int):
|
||||
def __new__(cls, *args, **kwargs):
|
||||
return super(IntWithCondOrUncond, cls).__new__(cls, *args, **kwargs)
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__()
|
||||
self.cond_or_uncond = None
|
||||
|
||||
class WrapperConsts:
|
||||
ACN = "ACN"
|
||||
VERSION = "version"
|
||||
ACN_OUTER_SAMPLE_WRAPPER_KEY = "ACN_outer_sample_wrapper"
|
||||
ACN_CREATE_SAMPLER_SAMPLE_WRAPPER = "create_outer_sample_wrapper"
|
||||
|
||||
|
||||
def get_properly_arranged_t2i_weights(initial_weights: list[float]):
|
||||
@@ -58,91 +144,76 @@ class ControlWeightType:
|
||||
UNIVERSAL = "universal"
|
||||
T2IADAPTER = "t2iadapter"
|
||||
CONTROLNET = "controlnet"
|
||||
CONTROLNETPLUSPLUS = "controlnet++"
|
||||
CONTROLLORA = "controllora"
|
||||
CONTROLLLLITE = "controllllite"
|
||||
SVD_CONTROLNET = "svd_controlnet"
|
||||
SPARSECTRL = "sparsectrl"
|
||||
CTRLORA = "ctrlora"
|
||||
|
||||
|
||||
class ControlWeights:
|
||||
def __init__(self, weight_type: str, base_multiplier: float=1.0,
|
||||
weights_input: list[float]=None, weights_middle: list[float]=None, weights_output: list[float]=None,
|
||||
weight_func: Callable=None, weight_mask: Tensor=None,
|
||||
uncond_multiplier=1.0, uncond_mask: Tensor=None,
|
||||
extras: dict[str]={}, disable_applied_to=False):
|
||||
def __init__(self, weight_type: str, base_multiplier: float=1.0, flip_weights: bool=False, weights: list[float]=None, weight_mask: Tensor=None,
|
||||
uncond_multiplier=1.0, uncond_mask: Tensor=None):
|
||||
self.weight_type = weight_type
|
||||
self.base_multiplier = base_multiplier
|
||||
self.weights_input = weights_input
|
||||
self.weights_middle = weights_middle
|
||||
self.weights_output = weights_output
|
||||
self.weight_func = weight_func
|
||||
self.flip_weights = flip_weights
|
||||
self.weights = weights
|
||||
if self.weights is not None and self.flip_weights:
|
||||
self.weights.reverse()
|
||||
self.weight_mask = weight_mask
|
||||
self.uncond_multiplier = float(uncond_multiplier)
|
||||
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.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
|
||||
def get(self, idx: int, default=1.0) -> Union[float, Tensor]:
|
||||
# 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:
|
||||
relevant_weights = list(reversed(relevant_weights))
|
||||
else:
|
||||
relevant_weights = self.weights_output
|
||||
if relevant_weights is None:
|
||||
return default * effective_mult
|
||||
elif idx >= len(relevant_weights):
|
||||
return default * effective_mult
|
||||
return relevant_weights[idx] * effective_mult
|
||||
if self.weights is not None:
|
||||
# this implies weights list is not aligning with expectations - will need to adjust code
|
||||
if idx >= len(self.weights):
|
||||
return default
|
||||
return self.weights[idx]
|
||||
return 1.0
|
||||
|
||||
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):
|
||||
return ControlWeights(weight_type=self.weight_type, base_multiplier=self.base_multiplier,
|
||||
weights_input=new_weights_input, weights_middle=new_weights_middle, weights_output=new_weights_output,
|
||||
weight_func=new_weight_func, weight_mask=self.weight_mask,
|
||||
uncond_multiplier=self.uncond_multiplier,
|
||||
extras=self.extras, disable_applied_to=self.disable_applied_to)
|
||||
def copy_with_new_weights(self, new_weights: list[float]):
|
||||
return ControlWeights(weight_type=self.weight_type, base_multiplier=self.base_multiplier, flip_weights=self.flip_weights,
|
||||
weights=new_weights, weight_mask=self.weight_mask, uncond_multiplier=self.uncond_multiplier)
|
||||
|
||||
@classmethod
|
||||
def default(cls, extras: dict[str]={}):
|
||||
return cls(ControlWeightType.DEFAULT, extras=extras)
|
||||
def default(cls):
|
||||
return cls(ControlWeightType.DEFAULT)
|
||||
|
||||
@classmethod
|
||||
def universal(cls, base_multiplier: float, uncond_multiplier: float=1.0, extras: dict[str]={}):
|
||||
return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, uncond_multiplier=uncond_multiplier, disable_applied_to=True, extras=extras)
|
||||
def universal(cls, base_multiplier: float, flip_weights: bool=False, uncond_multiplier: float=1.0):
|
||||
return cls(ControlWeightType.UNIVERSAL, base_multiplier=base_multiplier, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
|
||||
@classmethod
|
||||
def universal_mask(cls, weight_mask: Tensor, uncond_multiplier: float=1.0, extras: dict[str]={}):
|
||||
return cls(ControlWeightType.UNIVERSAL, weight_mask=weight_mask, uncond_multiplier=uncond_multiplier, disable_applied_to=True, extras=extras)
|
||||
def universal_mask(cls, weight_mask: Tensor, uncond_multiplier: float=1.0):
|
||||
return cls(ControlWeightType.UNIVERSAL, weight_mask=weight_mask, uncond_multiplier=uncond_multiplier)
|
||||
|
||||
@classmethod
|
||||
def t2iadapter(cls, weights_input: list[float]=None, uncond_multiplier: float=1.0, extras: dict[str]={}, disable_applied_to=False):
|
||||
return cls(ControlWeightType.T2IADAPTER, weights_input=weights_input, uncond_multiplier=uncond_multiplier, extras=extras, disable_applied_to=disable_applied_to)
|
||||
def t2iadapter(cls, weights: list[float]=None, flip_weights: bool=False, uncond_multiplier: float=1.0):
|
||||
if weights is None:
|
||||
weights = [1.0]*12
|
||||
return cls(ControlWeightType.T2IADAPTER, weights=weights,flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
|
||||
@classmethod
|
||||
def controlnet(cls, weights_output: list[float]=None, weights_middle: list[float]=None, weights_input: list[float]=None, uncond_multiplier: float=1.0, extras: dict[str]={}, disable_applied_to=False):
|
||||
return cls(ControlWeightType.CONTROLNET, weights_output=weights_output, weights_middle=weights_middle, weights_input=weights_input, uncond_multiplier=uncond_multiplier, extras=extras, disable_applied_to=disable_applied_to)
|
||||
def controlnet(cls, weights: list[float]=None, flip_weights: bool=False, uncond_multiplier: float=1.0):
|
||||
if weights is None:
|
||||
weights = [1.0]*13
|
||||
return cls(ControlWeightType.CONTROLNET, weights=weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
|
||||
@classmethod
|
||||
def controllora(cls, weights_output: list[float]=None, weights_middle: list[float]=None, weights_input: list[float]=None, uncond_multiplier: float=1.0, extras: dict[str]={}, disable_applied_to=False):
|
||||
return cls(ControlWeightType.CONTROLLORA, weights_output=weights_output, weights_middle=weights_middle, weights_input=weights_input, uncond_multiplier=uncond_multiplier, extras=extras, disable_applied_to=disable_applied_to)
|
||||
def controllora(cls, weights: list[float]=None, flip_weights: bool=False, uncond_multiplier: float=1.0):
|
||||
if weights is None:
|
||||
weights = [1.0]*10
|
||||
return cls(ControlWeightType.CONTROLLORA, weights=weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
|
||||
@classmethod
|
||||
def controllllite(cls, weights_output: list[float]=None, weights_middle: list[float]=None, weights_input: list[float]=None, uncond_multiplier: float=1.0, extras: dict[str]={}, disable_applied_to=False):
|
||||
return cls(ControlWeightType.CONTROLLLLITE, weights_output=weights_output, weights_middle=weights_middle, weights_input=weights_input, uncond_multiplier=uncond_multiplier, extras=extras, disable_applied_to=disable_applied_to)
|
||||
def controllllite(cls, weights: list[float]=None, flip_weights: bool=False, uncond_multiplier: float=1.0):
|
||||
if weights is None:
|
||||
# TODO: make this have a real value
|
||||
weights = [1.0]*200
|
||||
return cls(ControlWeightType.CONTROLLLLITE, weights=weights, flip_weights=flip_weights, uncond_multiplier=uncond_multiplier)
|
||||
|
||||
|
||||
class StrengthInterpolation:
|
||||
@@ -247,11 +318,6 @@ class TimestepKeyframe:
|
||||
def has_mask_hint(self):
|
||||
return self.mask_hint_orig is not None
|
||||
|
||||
def get_effective_guarantee_steps(self, max_sigma: torch.Tensor):
|
||||
'''If keyframe starts before current sampling range (max_sigma), treat as 0.'''
|
||||
if self.start_t > max_sigma:
|
||||
return 0
|
||||
return self.guarantee_steps
|
||||
|
||||
@staticmethod
|
||||
def default() -> 'TimestepKeyframe':
|
||||
@@ -260,10 +326,9 @@ class TimestepKeyframe:
|
||||
|
||||
# always maintain sorted state (by start_percent of TimestepKeyFrame)
|
||||
class TimestepKeyframeGroup:
|
||||
def __init__(self, add_default=True) -> None:
|
||||
def __init__(self) -> None:
|
||||
self.keyframes: list[TimestepKeyframe] = []
|
||||
if add_default:
|
||||
self.keyframes.append(TimestepKeyframe.default())
|
||||
self.keyframes.append(TimestepKeyframe.default())
|
||||
|
||||
def add(self, keyframe: TimestepKeyframe) -> None:
|
||||
# add to end of list, then sort
|
||||
@@ -289,7 +354,7 @@ class TimestepKeyframeGroup:
|
||||
return len(self.keyframes) == 0
|
||||
|
||||
def clone(self) -> 'TimestepKeyframeGroup':
|
||||
cloned = TimestepKeyframeGroup(add_default=False)
|
||||
cloned = TimestepKeyframeGroup()
|
||||
# already sorted, so don't use add function to make cloning quicker
|
||||
for tk in self.keyframes:
|
||||
cloned.keyframes.append(tk)
|
||||
@@ -304,7 +369,7 @@ class TimestepKeyframeGroup:
|
||||
|
||||
class AbstractPreprocWrapper:
|
||||
error_msg = "Invalid use of [InsertHere] output. The output of [InsertHere] preprocessor is NOT a usual image, but a latent pretending to be an image - you must connect the output directly to an Apply ControlNet node (advanced or otherwise). It cannot be used for anything else that accepts IMAGE input."
|
||||
def __init__(self, condhint):
|
||||
def __init__(self, condhint: Tensor):
|
||||
self.condhint = condhint
|
||||
|
||||
def movedim(self, *args, **kwargs):
|
||||
@@ -338,10 +403,8 @@ class AbstractPreprocWrapper:
|
||||
class disable_weight_init_clean_groupnorm(comfy.ops.disable_weight_init):
|
||||
class GroupNorm(comfy.ops.disable_weight_init.GroupNorm):
|
||||
def forward_comfy_cast_weights(self, input):
|
||||
weight, bias, offload_stream = comfy.ops.cast_bias_weight(self, input, offloadable=True)
|
||||
x = torch.nn.functional.group_norm(input, self.num_groups, weight, bias, self.eps)
|
||||
comfy.ops.uncast_bias_weight(self, weight, bias, offload_stream)
|
||||
return x
|
||||
weight, bias = comfy.ops.cast_bias_weight(self, input)
|
||||
return torch.nn.functional.group_norm(input, self.num_groups, weight, bias, self.eps)
|
||||
|
||||
def forward(self, input):
|
||||
if self.comfy_cast_weights:
|
||||
@@ -355,20 +418,11 @@ class manual_cast_clean_groupnorm(comfy.ops.manual_cast):
|
||||
|
||||
|
||||
# adapted from comfy/sample.py
|
||||
def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False, match_shape=False, flux_shape=None):
|
||||
def prepare_mask_batch(mask: Tensor, shape: Tensor, multiplier: int=1, match_dim1=False):
|
||||
mask = mask.clone()
|
||||
if flux_shape is not None:
|
||||
multiplier = multiplier * 0.5
|
||||
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(round(flux_shape[-2]*multiplier), round(flux_shape[-1]*multiplier)), mode="bilinear")
|
||||
mask = rearrange(mask, "b c h w -> b (h w) c")
|
||||
else:
|
||||
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(round(shape[-2]*multiplier), round(shape[-1]*multiplier)), mode="bilinear")
|
||||
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(shape[2]*multiplier, shape[3]*multiplier), mode="bilinear")
|
||||
if match_dim1:
|
||||
if match_shape and len(shape) < 4:
|
||||
raise Exception(f"match_dim1 cannot be True if shape is under 4 dims; was {len(shape)}.")
|
||||
mask = torch.cat([mask] * shape[1], dim=1)
|
||||
if match_shape and len(shape) == 3 and len(mask.shape) != 3:
|
||||
mask = mask.squeeze(1)
|
||||
return mask
|
||||
|
||||
|
||||
@@ -470,25 +524,15 @@ def get_sorted_list_via_attr(objects: list, attr: str) -> list:
|
||||
return sorted_list
|
||||
|
||||
|
||||
# DFS Search for Torch.nn.Module, Written by Lvmin
|
||||
def torch_dfs(model: torch.nn.Module):
|
||||
result = [model]
|
||||
for child in model.children():
|
||||
result += torch_dfs(child)
|
||||
return result
|
||||
|
||||
|
||||
class WeightTypeException(TypeError):
|
||||
"Raised when weight not compatible with AdvancedControlBase object"
|
||||
pass
|
||||
|
||||
|
||||
class AdvancedControlBase:
|
||||
ACN_VERSION = CURRENT_WRAPPER_VERSION
|
||||
|
||||
def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights, require_vae=False, allow_condhint_latents=False):
|
||||
def __init__(self, base: ControlBase, timestep_keyframes: TimestepKeyframeGroup, weights_default: ControlWeights, require_model=False):
|
||||
self.base = base
|
||||
self.compatible_weights = [ControlWeightType.UNIVERSAL, ControlWeightType.DEFAULT]
|
||||
self.compatible_weights = [ControlWeightType.UNIVERSAL]
|
||||
self.add_compatible_weight(weights_default.weight_type)
|
||||
# mask for which parts of controlnet output to keep
|
||||
self.mask_cond_hint_original = None
|
||||
@@ -501,11 +545,9 @@ class AdvancedControlBase:
|
||||
self.full_latent_length = 0
|
||||
self.context_length = 0
|
||||
# timesteps
|
||||
self.t: float = None
|
||||
self.prev_t: float = None
|
||||
self.batched_number: int = None
|
||||
self.t: Tensor = None
|
||||
self.batched_number: Union[int, IntWithCondOrUncond] = None
|
||||
self.batch_size: int = 0
|
||||
self.cond_or_uncond: list[int] = None
|
||||
# weights + override
|
||||
self.weights: ControlWeights = None
|
||||
self.weights_default: ControlWeights = weights_default
|
||||
@@ -521,18 +563,13 @@ class AdvancedControlBase:
|
||||
self.pre_run = self.pre_run_inject
|
||||
self.cleanup = self.cleanup_inject
|
||||
self.set_previous_controlnet = self.set_previous_controlnet_inject
|
||||
self.set_cond_hint = self.set_cond_hint_inject
|
||||
# vae to store
|
||||
self.adv_vae = None
|
||||
self.mult_by_ratio_when_vae = True
|
||||
# compression ratio stuff
|
||||
self.real_compression_ratio = None
|
||||
# require model/vae to be passed into Apply Advanced ControlNet 🛂🅐🅒🅝 node
|
||||
self.require_vae = require_vae
|
||||
self.allow_condhint_latents = allow_condhint_latents
|
||||
self.postpone_condhint_latents_check = False
|
||||
# require model to be passed into Apply Advanced ControlNet 🛂🅐🅒🅝 node
|
||||
self.require_model = require_model
|
||||
# disarm - when set to False, used to force usage of Apply Advanced ControlNet 🛂🅐🅒🅝 node (which will set it to True)
|
||||
self.disarmed = True
|
||||
self.disarmed = not require_model
|
||||
|
||||
def patch_model(self, model: ModelPatcher):
|
||||
pass
|
||||
|
||||
def add_compatible_weight(self, control_weight_type: str):
|
||||
self.compatible_weights.append(control_weight_type)
|
||||
@@ -548,7 +585,7 @@ class AdvancedControlBase:
|
||||
else:
|
||||
for tk in self.timestep_keyframes.keyframes:
|
||||
if tk.has_control_weights() and tk.control_weights.weight_type not in self.compatible_weights:
|
||||
msg = f"Weight on Timestep Keyframe with start_percent={tk.start_percent} is type " + \
|
||||
msg = f"Weight on Timestep Keyframe with start_percent={tk.start_percent} is type" + \
|
||||
f"{tk.control_weights.weight_type}, but loaded {type(self).__name__} only supports {self.compatible_weights} weights."
|
||||
raise WeightTypeException(msg)
|
||||
|
||||
@@ -561,17 +598,15 @@ class AdvancedControlBase:
|
||||
self.weights = None
|
||||
self.latent_keyframes = None
|
||||
|
||||
def prepare_current_timestep(self, t: Tensor, transformer_options: dict[str, torch.Tensor]):
|
||||
def prepare_current_timestep(self, t: Tensor, batched_number: int):
|
||||
self.t = float(t[0])
|
||||
# check if t has changed (otherwise do nothing, as step already accounted for)
|
||||
if self.t == self.prev_t:
|
||||
return
|
||||
self.batched_number = batched_number
|
||||
self.batch_size = len(t)
|
||||
# get current step percent
|
||||
curr_t: float = self.t
|
||||
prev_index = self._current_timestep_index
|
||||
max_sigma = torch.max(transformer_options.get("sample_sigmas", BIGMAX_TENSOR))
|
||||
# if met guaranteed steps (or no current keyframe), look for next keyframe in case need to switch
|
||||
if self._current_timestep_keyframe is None or self._current_used_steps >= self._current_timestep_keyframe.get_effective_guarantee_steps(max_sigma):
|
||||
if self._current_timestep_keyframe is None or self._current_used_steps >= self._current_timestep_keyframe.guarantee_steps:
|
||||
# if has next index, loop through and see if need to switch
|
||||
if self.timestep_keyframes.has_index(self._current_timestep_index+1):
|
||||
for i in range(self._current_timestep_index+1, len(self.timestep_keyframes)):
|
||||
@@ -597,13 +632,12 @@ class AdvancedControlBase:
|
||||
del self.tk_mask_cond_hint_original
|
||||
self.tk_mask_cond_hint_original = None
|
||||
# if guarantee_steps greater than zero, stop searching for other keyframes
|
||||
if self._current_timestep_keyframe.get_effective_guarantee_steps(max_sigma) > 0:
|
||||
if self._current_timestep_keyframe.guarantee_steps > 0:
|
||||
break
|
||||
# if eval_tk is outside of percent range, stop looking further
|
||||
else:
|
||||
break
|
||||
# update prev_t
|
||||
self.prev_t = self.t
|
||||
|
||||
# update steps current keyframe is used
|
||||
self._current_used_steps += 1
|
||||
# if index changed, apply overrides
|
||||
@@ -618,7 +652,7 @@ class AdvancedControlBase:
|
||||
self.prepare_weights()
|
||||
|
||||
def prepare_weights(self):
|
||||
if self.weights is None:
|
||||
if self.weights is None or self.weights.weight_type == ControlWeightType.DEFAULT:
|
||||
self.weights = self.weights_default
|
||||
elif self.weights.weight_type == ControlWeightType.UNIVERSAL:
|
||||
# if universal and weight_mask present, no need to convert
|
||||
@@ -633,23 +667,6 @@ class AdvancedControlBase:
|
||||
self.mask_cond_hint_original = mask_hint
|
||||
return self
|
||||
|
||||
def set_cond_hint_inject(self, *args, **kwargs):
|
||||
to_return = self.base.set_cond_hint(*args, **kwargs)
|
||||
# if vae required, look in args and kwargs for it
|
||||
if self.require_vae:
|
||||
# check args first, as that's the default way vae param is used in ComfyUI
|
||||
for arg in args:
|
||||
if isinstance(arg, VAE):
|
||||
self.adv_vae = arg
|
||||
self.vae = arg
|
||||
break
|
||||
# if not in args, check kwargs now
|
||||
if self.adv_vae is None:
|
||||
if 'vae' in kwargs:
|
||||
self.adv_vae = kwargs['vae']
|
||||
self.vae = kwargs['vae']
|
||||
return to_return
|
||||
|
||||
def pre_run_inject(self, model, percent_to_timestep_function):
|
||||
self.base.pre_run(model, percent_to_timestep_function)
|
||||
self.pre_run_advanced(model, percent_to_timestep_function)
|
||||
@@ -658,9 +675,6 @@ class AdvancedControlBase:
|
||||
# for each timestep keyframe, calculate the start_t
|
||||
for tk in self.timestep_keyframes.keyframes:
|
||||
tk.start_t = percent_to_timestep_function(tk.start_percent)
|
||||
# set real_compression_ratio to compression_ratio
|
||||
if hasattr(self, "compression_ratio"):
|
||||
self.real_compression_ratio = self.compression_ratio
|
||||
# clear variables
|
||||
self.cleanup_advanced()
|
||||
|
||||
@@ -681,49 +695,34 @@ class AdvancedControlBase:
|
||||
return False
|
||||
return True
|
||||
|
||||
def get_control_inject(self, x_noisy, t, cond, batched_number, transformer_options: dict):
|
||||
self.batched_number = batched_number
|
||||
self.batch_size = len(t)
|
||||
self.cond_or_uncond = transformer_options.get("cond_or_uncond", None)
|
||||
# fill out ad_param-related fields, if present
|
||||
if "ad_params" in transformer_options:
|
||||
self.sub_idxs = transformer_options["ad_params"]["sub_idxs"]
|
||||
self.full_latent_length = transformer_options["ad_params"]["full_length"]
|
||||
self.context_length = transformer_options["ad_params"]["context_length"]
|
||||
def get_control_inject(self, x_noisy, t, cond, batched_number):
|
||||
# prepare timestep and everything related
|
||||
self.prepare_current_timestep(t=t, transformer_options=transformer_options)
|
||||
self.prepare_current_timestep(t=t, batched_number=batched_number)
|
||||
# if should not perform any actions for the controlnet, exit without doing any work
|
||||
if self.strength == 0.0 or self._current_timestep_keyframe.strength == 0.0:
|
||||
return self.default_control_actions(x_noisy, t, cond, batched_number, transformer_options)
|
||||
return self.default_control_actions(x_noisy, t, cond, batched_number)
|
||||
# otherwise, perform normal function
|
||||
return self.get_control_advanced(x_noisy, t, cond, batched_number, transformer_options)
|
||||
return self.get_control_advanced(x_noisy, t, cond, batched_number)
|
||||
|
||||
def get_control_advanced(self, x_noisy, t, cond, batched_number, transformer_options):
|
||||
return self.default_control_actions(x_noisy, t, cond, batched_number, transformer_options)
|
||||
def get_control_advanced(self, x_noisy, t, cond, batched_number):
|
||||
return self.default_control_actions(x_noisy, t, cond, batched_number)
|
||||
|
||||
def default_control_actions(self, x_noisy, t, cond, batched_number, transformer_options):
|
||||
def default_control_actions(self, x_noisy, t, cond, batched_number):
|
||||
control_prev = None
|
||||
if self.previous_controlnet is not None:
|
||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number, transformer_options)
|
||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
|
||||
return control_prev
|
||||
|
||||
def calc_weight(self, idx: int, x: Tensor, control: dict[str, list[Tensor]], key: str) -> Union[float, Tensor]:
|
||||
def calc_weight(self, idx: int, x: Tensor, layers: int) -> Union[float, Tensor]:
|
||||
if self.weights.weight_mask is not None:
|
||||
# prepare weight mask
|
||||
self.prepare_weight_mask_cond_hint(x, self.batched_number)
|
||||
# adjust mask for current layer and return
|
||||
return torch.pow(self.weight_mask_cond_hint, self.get_calc_pow(idx=idx, control=control, key=key))
|
||||
return self.weights.get(idx=idx, control=control, key=key)
|
||||
return torch.pow(self.weight_mask_cond_hint, self.get_calc_pow(idx=idx, layers=layers))
|
||||
return self.weights.get(idx=idx)
|
||||
|
||||
def get_calc_pow(self, idx: int, control: dict[str, list[Tensor]], key: str) -> int:
|
||||
if key == "middle":
|
||||
return 0
|
||||
else:
|
||||
c_len = len(control[key])
|
||||
real_idx = c_len-idx
|
||||
if key == "input":
|
||||
real_idx = c_len - real_idx + 1
|
||||
return real_idx
|
||||
def get_calc_pow(self, idx: int, layers: int) -> int:
|
||||
return (layers-1)-idx
|
||||
|
||||
def calc_latent_keyframe_mults(self, x: Tensor, batched_number: int) -> Tensor:
|
||||
# apply strengths, and get batch indeces to null out
|
||||
@@ -773,11 +772,12 @@ class AdvancedControlBase:
|
||||
final_tensor = final_tensor.unsqueeze(-1)
|
||||
return final_tensor
|
||||
|
||||
def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int, flux_shape: tuple=None):
|
||||
def apply_advanced_strengths_and_masks(self, x: Tensor, batched_number: int):
|
||||
# handle weight's uncond_multiplier, if applicable
|
||||
if self.weights.has_uncond_multiplier:
|
||||
cond_or_uncond = self.batched_number.cond_or_uncond
|
||||
actual_length = x.size(0) // batched_number
|
||||
for idx, cond_type in enumerate(self.cond_or_uncond):
|
||||
for idx, cond_type in enumerate(cond_or_uncond):
|
||||
# if uncond, set to weight's uncond_multiplier
|
||||
if cond_type == 1:
|
||||
x[actual_length*idx:actual_length*(idx+1)] *= self.weights.uncond_multiplier
|
||||
@@ -788,41 +788,50 @@ class AdvancedControlBase:
|
||||
x[:] = x[:] * self.calc_latent_keyframe_mults(x=x, batched_number=batched_number)
|
||||
# apply masks, resizing mask to required dims
|
||||
if self.mask_cond_hint is not None:
|
||||
masks = prepare_mask_batch(self.mask_cond_hint, x.shape, match_shape=True, flux_shape=flux_shape)
|
||||
masks = prepare_mask_batch(self.mask_cond_hint, x.shape)
|
||||
x[:] = x[:] * masks
|
||||
if self.tk_mask_cond_hint is not None:
|
||||
masks = prepare_mask_batch(self.tk_mask_cond_hint, x.shape, match_shape=True, flux_shape=flux_shape)
|
||||
masks = prepare_mask_batch(self.tk_mask_cond_hint, x.shape)
|
||||
x[:] = x[:] * masks
|
||||
# apply timestep keyframe strengths
|
||||
if self._current_timestep_keyframe.strength != 1.0:
|
||||
x[:] *= self._current_timestep_keyframe.strength
|
||||
|
||||
def control_merge_inject(self: 'AdvancedControlBase', control: dict[str, list[Tensor]], control_prev: dict, output_dtype):
|
||||
def control_merge_inject(self: 'AdvancedControlBase', control_input, control_output, control_prev, output_dtype):
|
||||
out = {'input':[], 'middle':[], 'output': []}
|
||||
|
||||
for key in control:
|
||||
control_output = control[key]
|
||||
applied_to = set()
|
||||
if control_input is not None:
|
||||
for i in range(len(control_input)):
|
||||
key = 'input'
|
||||
x = control_input[i]
|
||||
if x is not None:
|
||||
self.apply_advanced_strengths_and_masks(x, self.batched_number)
|
||||
|
||||
x *= self.strength * self.calc_weight(i, x, len(control_input))
|
||||
if x.dtype != output_dtype:
|
||||
x = x.to(output_dtype)
|
||||
out[key].insert(0, x)
|
||||
|
||||
if control_output is not None:
|
||||
for i in range(len(control_output)):
|
||||
if i == (len(control_output) - 1):
|
||||
key = 'middle'
|
||||
index = 0
|
||||
else:
|
||||
key = 'output'
|
||||
index = i
|
||||
x = control_output[i]
|
||||
if x is not None:
|
||||
self.apply_advanced_strengths_and_masks(x, self.batched_number)
|
||||
|
||||
if self.global_average_pooling:
|
||||
x = torch.mean(x, dim=(2, 3), keepdim=True).repeat(1, 1, x.shape[2], x.shape[3])
|
||||
|
||||
# if should disable applied_to optimization, clone the weight if in applied_to
|
||||
if self.weights.disable_applied_to and x in applied_to:
|
||||
x = x.clone()
|
||||
|
||||
if x not in applied_to: #memory saving strategy, allow shared tensors and only apply strength to shared tensors once
|
||||
applied_to.add(x)
|
||||
self.apply_advanced_strengths_and_masks(x, self.batched_number)
|
||||
x *= self.strength * self.calc_weight(i, x, control, key)
|
||||
|
||||
if output_dtype is not None and x.dtype != output_dtype:
|
||||
x *= self.strength * self.calc_weight(i, x, len(control_output))
|
||||
if x.dtype != output_dtype:
|
||||
x = x.to(output_dtype)
|
||||
|
||||
out[key].append(x)
|
||||
|
||||
if control_prev is not None:
|
||||
for x in ['input', 'middle', 'output']:
|
||||
o = out[x]
|
||||
@@ -837,7 +846,7 @@ class AdvancedControlBase:
|
||||
if o[i].shape[0] < prev_val.shape[0]:
|
||||
o[i] = prev_val + o[i]
|
||||
else:
|
||||
o[i] = prev_val + o[i] # TODO from base ComfyUI: change back to inplace add if shared tensors stop being an issue
|
||||
o[i] += prev_val
|
||||
return out
|
||||
|
||||
def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False):
|
||||
@@ -860,7 +869,7 @@ class AdvancedControlBase:
|
||||
del out_mask
|
||||
# TODO: perform upscale on only the sub_idxs masks at a time instead of all to conserve RAM
|
||||
# resize mask and match batch count
|
||||
out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier, match_shape=True)
|
||||
out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier)
|
||||
actual_latent_length = x_noisy.shape[0] // batched_number
|
||||
out_mask = extend_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length)
|
||||
if self.sub_idxs is not None:
|
||||
@@ -871,7 +880,7 @@ class AdvancedControlBase:
|
||||
# default dtype to be same as x_noisy
|
||||
if dtype is None:
|
||||
dtype = x_noisy.dtype
|
||||
setattr(self, attr_name, out_mask.to(dtype=dtype).to(x_noisy.device))
|
||||
setattr(self, attr_name, out_mask.to(dtype=dtype).to(self.device))
|
||||
del out_mask
|
||||
|
||||
def _reset_attr(self, attr_name, new_value=None):
|
||||
@@ -888,14 +897,10 @@ class AdvancedControlBase:
|
||||
self.full_latent_length = 0
|
||||
self.context_length = 0
|
||||
self.t = None
|
||||
self.prev_t = None
|
||||
self.batched_number = None
|
||||
self.batch_size = 0
|
||||
self.weights = None
|
||||
self.latent_keyframes = None
|
||||
# set effective_compression_ratio to compression_ratio
|
||||
if hasattr(self, "compression_ratio"):
|
||||
self.real_compression_ratio = self.compression_ratio
|
||||
# timestep stuff
|
||||
self._current_timestep_keyframe = None
|
||||
self._current_timestep_index = -1
|
||||
@@ -918,9 +923,4 @@ class AdvancedControlBase:
|
||||
copied.mask_cond_hint_original = self.mask_cond_hint_original
|
||||
copied.weights_override = self.weights_override
|
||||
copied.latent_keyframe_override = self.latent_keyframe_override
|
||||
copied.adv_vae = self.adv_vae
|
||||
copied.mult_by_ratio_when_vae = self.mult_by_ratio_when_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
|
||||
|
||||
+2
-3
@@ -1,8 +1,8 @@
|
||||
[project]
|
||||
name = "comfyui-advanced-controlnet"
|
||||
description = "Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks."
|
||||
version = "1.5.8"
|
||||
license = { file = "LICENSE" }
|
||||
version = "1.0.2"
|
||||
license = "LICENSE"
|
||||
dependencies = []
|
||||
|
||||
[project.urls]
|
||||
@@ -13,4 +13,3 @@ Repository = "https://github.com/Kosinkadink/ComfyUI-Advanced-ControlNet"
|
||||
PublisherId = "kosinkadink"
|
||||
DisplayName = "ComfyUI-Advanced-ControlNet"
|
||||
Icon = ""
|
||||
requires-comfyui = ">=0.3.68"
|
||||
|
||||
@@ -1,53 +0,0 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
|
||||
function addResizeHook(node, padding, useOldMin=false) {
|
||||
let origOnCreated = node.onNodeCreated
|
||||
node.onNodeCreated = function() {
|
||||
let r = origOnCreated?.apply(this, arguments)
|
||||
let size = this.computeSize();
|
||||
size[0] += padding || 0;
|
||||
if (useOldMin) {
|
||||
//equal to LiteGraph.NODE_WIDTH*1.5*1.5
|
||||
size[0] = Math.max(size[0], 315)
|
||||
}
|
||||
this.setSize(size);
|
||||
return r
|
||||
}
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "AdvancedControlNet.autosize",
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
//since python_module is based off folder path,
|
||||
//it could be changed by users and should only be used as fallback
|
||||
if (nodeData?.name?.startsWith("ACN_")
|
||||
|| nodeData.python_module == 'custom_nodes.ComfyUI-Advanced-ControlNet') {
|
||||
if (nodeData?.input?.hidden?.autosize) {
|
||||
addResizeHook(nodeType.prototype, nodeData.input.hidden.autosize[1]?.padding)
|
||||
} else if (!nodeData?.input?.optional?.autosize) {
|
||||
addResizeHook(nodeType.prototype, 0, true)
|
||||
}
|
||||
}
|
||||
},
|
||||
async getCustomWidgets() {
|
||||
return {
|
||||
ACNAUTOSIZE(node, inputName, inputData) {
|
||||
let w = {
|
||||
name : inputName,
|
||||
type : "ACN.AUTOSIZE",
|
||||
value : "",
|
||||
options : {"serialize": false},
|
||||
computeSize : function(width) {
|
||||
return [0, -4];
|
||||
}
|
||||
}
|
||||
if (!node.widgets) {
|
||||
node.widgets = []
|
||||
}
|
||||
node.widgets.push(w)
|
||||
addResizeHook(node, inputData[1].padding);
|
||||
return w;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
@@ -1,293 +0,0 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
|
||||
function chainCallback(object, property, callback) {
|
||||
if (object == undefined) {
|
||||
//This should not happen.
|
||||
console.error("Tried to add callback to non-existant object")
|
||||
return;
|
||||
}
|
||||
if (property in object && object[property]) {
|
||||
const callback_orig = object[property]
|
||||
object[property] = function () {
|
||||
const r = callback_orig.apply(this, arguments);
|
||||
callback.apply(this, arguments);
|
||||
return r
|
||||
};
|
||||
} else {
|
||||
object[property] = callback;
|
||||
}
|
||||
}
|
||||
var helpDOM;
|
||||
function initHelpDOM() {
|
||||
let parentDOM = document.createElement("div");
|
||||
document.body.appendChild(parentDOM)
|
||||
parentDOM.appendChild(helpDOM)
|
||||
helpDOM.className = "litegraph";
|
||||
let scrollbarStyle = document.createElement('style');
|
||||
scrollbarStyle.innerHTML = `
|
||||
<style id="scroll-properties">
|
||||
* {
|
||||
scrollbar-width: 6px;
|
||||
scrollbar-color: #0003 #0000;
|
||||
}
|
||||
::-webkit-scrollbar {
|
||||
background: transparent;
|
||||
width: 6px;
|
||||
}
|
||||
::-webkit-scrollbar-thumb {
|
||||
background: #0005;
|
||||
border-radius: 20px
|
||||
}
|
||||
::-webkit-scrollbar-button {
|
||||
display: none;
|
||||
}
|
||||
.VHS_loopedvideo::-webkit-media-controls-mute-button {
|
||||
display:none;
|
||||
}
|
||||
.VHS_loopedvideo::-webkit-media-controls-fullscreen-button {
|
||||
display:none;
|
||||
}
|
||||
</style>
|
||||
`
|
||||
parentDOM.appendChild(scrollbarStyle)
|
||||
chainCallback(app.canvas, "onDrawForeground", function (ctx, visible_rect){
|
||||
let n = helpDOM.node
|
||||
if (!n || !n?.graph) {
|
||||
parentDOM.style['left'] = '-5000px'
|
||||
return
|
||||
}
|
||||
//draw : function(ctx, node, widgetWidth, widgetY, height) {
|
||||
//update widget position, even if off screen
|
||||
const transform = ctx.getTransform();
|
||||
const scale = app.canvas.ds.scale;//gets the litegraph zoom
|
||||
//calculate coordinates with account for browser zoom
|
||||
const bcr = app.canvas.canvas.getBoundingClientRect()
|
||||
const x = transform.e*scale/transform.a + bcr.x;
|
||||
const y = transform.f*scale/transform.a + bcr.y;
|
||||
//TODO: text reflows at low zoom. investigate alternatives
|
||||
Object.assign(parentDOM.style, {
|
||||
left: (x+(n.pos[0] + n.size[0]+15)*scale) + "px",
|
||||
top: (y+(n.pos[1]-LiteGraph.NODE_TITLE_HEIGHT)*scale) + "px",
|
||||
width: "400px",
|
||||
minHeight: "100px",
|
||||
maxHeight: "600px",
|
||||
overflowY: 'scroll',
|
||||
transformOrigin: '0 0',
|
||||
transform: 'scale(' + scale + ',' + scale +')',
|
||||
fontSize: '18px',
|
||||
backgroundColor: LiteGraph.NODE_DEFAULT_BGCOLOR,
|
||||
boxShadow: '0 0 10px black',
|
||||
borderRadius: '4px',
|
||||
padding: '3px',
|
||||
zIndex: 3,
|
||||
position: "absolute",
|
||||
display: 'inline',
|
||||
});
|
||||
});
|
||||
function setCollapse(el, doCollapse) {
|
||||
if (doCollapse) {
|
||||
el.children[0].children[0].innerHTML = '+'
|
||||
Object.assign(el.children[1].style, {
|
||||
color: '#CCC',
|
||||
overflowX: 'hidden',
|
||||
width: '0px',
|
||||
minWidth: 'calc(100% - 20px)',
|
||||
textOverflow: 'ellipsis',
|
||||
whiteSpace: 'nowrap',
|
||||
})
|
||||
for (let child of el.children[1].children) {
|
||||
if (child.style.display != 'none'){
|
||||
child.origDisplay = child.style.display
|
||||
}
|
||||
child.style.display = 'none'
|
||||
}
|
||||
} else {
|
||||
el.children[0].children[0].innerHTML = '-'
|
||||
Object.assign(el.children[1].style, {
|
||||
color: '',
|
||||
overflowX: '',
|
||||
width: '100%',
|
||||
minWidth: '',
|
||||
textOverflow: '',
|
||||
whiteSpace: '',
|
||||
})
|
||||
for (let child of el.children[1].children) {
|
||||
child.style.display = child.origDisplay
|
||||
}
|
||||
}
|
||||
}
|
||||
helpDOM.collapseOnClick = function() {
|
||||
let doCollapse = this.children[0].innerHTML == '-'
|
||||
setCollapse(this.parentElement, doCollapse)
|
||||
}
|
||||
helpDOM.selectHelp = function(name, value) {
|
||||
//attempt to navigate to name in help
|
||||
function collapseUnlessMatch(items,t) {
|
||||
var match = items.querySelector('[vhs_title="' + t + '"]')
|
||||
if (!match) {
|
||||
for (let i of items.children) {
|
||||
if (i.innerHTML.slice(0,t.length+5).includes(t)) {
|
||||
match = i
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
if (!match) {
|
||||
return null
|
||||
}
|
||||
//For longer documentation items with fewer collapsable elements,
|
||||
//scroll to make sure the entirety of the selected item is visible
|
||||
//This has the unfortunate side effect of trying to scroll the main
|
||||
//window if the documentation windows is forcibly offscreen,
|
||||
//but it's easy to simply scroll the main window back and seems to
|
||||
//have no visual side effects
|
||||
match.scrollIntoView(false)
|
||||
window.scrollTo(0,0)
|
||||
for (let i of items.querySelectorAll('.VHS_collapse')) {
|
||||
if (i.contains(match)) {
|
||||
setCollapse(i, false)
|
||||
} else {
|
||||
setCollapse(i, true)
|
||||
}
|
||||
}
|
||||
return match
|
||||
}
|
||||
let target = collapseUnlessMatch(helpDOM, name)
|
||||
if (target && value) {
|
||||
collapseUnlessMatch(target, value)
|
||||
}
|
||||
}
|
||||
|
||||
helpDOM.addHelp = function(node, nodeType, description) {
|
||||
if (!description) {
|
||||
return
|
||||
}
|
||||
//Pad computed size for the clickable question mark
|
||||
let originalComputeSize = node.computeSize
|
||||
node.computeSize = function() {
|
||||
let size = originalComputeSize.apply(this, arguments)
|
||||
if (!this.title) {
|
||||
return size
|
||||
}
|
||||
let title_width = this.title.length * 0.6 * LiteGraph.NODE_TEXT_SIZE
|
||||
size[0] = Math.max(size[0], title_width + LiteGraph.NODE_TITLE_HEIGHT)
|
||||
return size
|
||||
}
|
||||
|
||||
node.description = description
|
||||
chainCallback(node, "onDrawForeground", function (ctx) {
|
||||
//draw question mark
|
||||
ctx.save()
|
||||
ctx.font = 'bold 20px Arial'
|
||||
ctx.fillText("?", this.size[0]-17, -8)
|
||||
ctx.restore()
|
||||
})
|
||||
chainCallback(node, "onMouseDown", function (e, pos, canvas) {
|
||||
//On click would be preferred, but this'll be good enough
|
||||
if (pos[1] < 0 && pos[0] + LiteGraph.NODE_TITLE_HEIGHT > this.size[0]) {
|
||||
//corner question mark clicked
|
||||
if (helpDOM.node == this) {
|
||||
helpDOM.node = undefined
|
||||
} else {
|
||||
helpDOM.node = this;
|
||||
helpDOM.innerHTML = this.description || "no help provided ".repeat(20)
|
||||
for (let e of helpDOM.querySelectorAll('.VHS_collapse')) {
|
||||
e.children[0].onclick = helpDOM.collapseOnClick
|
||||
e.children[0].style.cursor = 'pointer'
|
||||
}
|
||||
for (let e of helpDOM.querySelectorAll('.VHS_precollapse')) {
|
||||
setCollapse(e, true)
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
})
|
||||
let timeout = null
|
||||
chainCallback(node, "onMouseMove", function (e, pos, canvas) {
|
||||
if (timeout) {
|
||||
clearTimeout(timeout)
|
||||
timeout = null
|
||||
}
|
||||
if (helpDOM.node != this) {
|
||||
return
|
||||
}
|
||||
timeout = setTimeout(() => {
|
||||
let n = this
|
||||
if (pos[0] > 0 && pos[0] < n.size[0]
|
||||
&& pos[1] > 0 && pos[1] < n.size[1]) {
|
||||
//TODO: provide help specific to element clicked
|
||||
let inputRows = Math.max(n.inputs.length, n.outputs.length)
|
||||
if (pos[1] < LiteGraph.NODE_SLOT_HEIGHT * inputRows) {
|
||||
let row = Math.floor((pos[1] - 7) / LiteGraph.NODE_SLOT_HEIGHT)
|
||||
if (pos[0] < n.size[0]/2) {
|
||||
if (row < n.inputs.length) {
|
||||
helpDOM.selectHelp(n.inputs[row].name)
|
||||
}
|
||||
} else {
|
||||
if (row < n.outputs.length) {
|
||||
helpDOM.selectHelp(n.outputs[row].name)
|
||||
}
|
||||
}
|
||||
} else {
|
||||
//probably widget, but widgets have variable height.
|
||||
let basey = LiteGraph.NODE_SLOT_HEIGHT * inputRows + 6
|
||||
for (let w of n.widgets) {
|
||||
if (w.y) {
|
||||
basey = w.y
|
||||
}
|
||||
let wheight = LiteGraph.NODE_WIDGET_HEIGHT+4
|
||||
if (w.computeSize) {
|
||||
wheight = w.computeSize(n.size[0])[1]
|
||||
}
|
||||
if (pos[1] < basey + wheight) {
|
||||
helpDOM.selectHelp(w.name, w.value)
|
||||
break
|
||||
}
|
||||
basey += wheight
|
||||
}
|
||||
}
|
||||
}
|
||||
}, 500)
|
||||
})
|
||||
chainCallback(node, "onMouseLeave", function (e, pos, canvas) {
|
||||
if (timeout) {
|
||||
clearTimeout(timeout)
|
||||
timeout = null
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
app.registerExtension({
|
||||
name: "AdvancedControlNet.documentation",
|
||||
async init() {
|
||||
if (app.VHSHelp) {
|
||||
helpDOM = app.VHSHelp
|
||||
} else {
|
||||
helpDOM = document.createElement("div");
|
||||
initHelpDOM()
|
||||
app.VHSHelp = helpDOM
|
||||
}
|
||||
},
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
// NOTE: May need manual adjusting for the few non-namespaced nodes
|
||||
if(nodeData?.name?.startsWith("ACN_") && nodeData.description) {
|
||||
let description = nodeData.description
|
||||
let el = document.createElement("div")
|
||||
el.innerHTML = description
|
||||
if (!el.children.length) {
|
||||
//Is plaintext. Do minor convenience formatting
|
||||
let chunks = description.split('\n')
|
||||
nodeData.description = chunks[0]
|
||||
description = chunks.join('<br>')
|
||||
} else {
|
||||
nodeData.description = el.querySelector('#VHS_shortdesc')?.innerHTML || el.children[1]?.firstChild?.innerHTML
|
||||
}
|
||||
chainCallback(nodeType.prototype, "onNodeCreated", function () {
|
||||
helpDOM.addHelp(this, nodeType, description)
|
||||
})
|
||||
}
|
||||
},
|
||||
});
|
||||
Reference in New Issue
Block a user