Merge PR #92 from Kosinkadink/develop - Reference support
Reference ControlNet support (reference_attn)
This commit is contained in:
@@ -1,5 +1,5 @@
|
||||
# ComfyUI-Advanced-ControlNet
|
||||
Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks. The ControlNet nodes here fully support sliding context sampling, like the one used in the [ComfyUI-AnimateDiff-Evolved](https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved) nodes. Currently supports ControlNets, T2IAdapters, ControlLoRAs, ControlLLLite, SparseCtrls, and SVD-ControlNets.
|
||||
Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks. The ControlNet nodes here fully support sliding context sampling, like the one used in the [ComfyUI-AnimateDiff-Evolved](https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved) nodes. Currently supports ControlNets, T2IAdapters, ControlLoRAs, ControlLLLite, SparseCtrls, SVD-ControlNets, and Reference.
|
||||
|
||||
Custom weights allow replication of the "My prompt is more important" feature of Auto1111's sd-webui ControlNet extension.
|
||||
|
||||
@@ -8,12 +8,14 @@ ControlNet preprocessors are available through [comfyui_controlnet_aux](https://
|
||||
## Features
|
||||
- Timestep and latent strength scheduling
|
||||
- Attention masks
|
||||
- Soft weights to replicate "My prompt is more important" feature from sd-webui ControlNet extension, and also change the scaling.
|
||||
- ControlNet, T2IAdapter, and ControlLoRA support for sliding context windows.
|
||||
- Soft weights to replicate "My prompt is more important" feature from sd-webui ControlNet extension, and also change the scaling
|
||||
- ControlNet, T2IAdapter, and ControlLoRA support for sliding context windows
|
||||
- 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)
|
||||
- Reference support
|
||||
- Currently, only ```reference_attn``` is exposed (equivalent of reference_only in Auto1111), reference_adain (and +attn) under construction. ```style_fidelity``` and ```ref_weight``` are equivalent to style_fidelity and control_weight in Auto1111, respectively, and strength of the Apply ControlNet is the balance between ref-influenced result and no-ref result.
|
||||
|
||||
## Table of Contents:
|
||||
- [Scheduling Explanation](#scheduling-explanation)
|
||||
|
||||
@@ -349,60 +349,6 @@ class SparseCtrlAdvanced(ControlNetAdvanced):
|
||||
return c
|
||||
|
||||
|
||||
class ReferenceAdvanced(ControlBase, AdvancedControlBase):
|
||||
def __init__(self, timestep_keyframes: TimestepKeyframeGroup, device=None):
|
||||
super().__init__(device)
|
||||
AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite(), require_model=True)
|
||||
# TODO: save attn patches here
|
||||
|
||||
def patch_model(self, model: ModelPatcher):
|
||||
# TODO: do model patching here
|
||||
pass
|
||||
|
||||
def pre_run_advanced(self, *args, **kwargs):
|
||||
AdvancedControlBase.pre_run_advanced(self, *args, **kwargs)
|
||||
# TODO: set control on patches
|
||||
|
||||
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int):
|
||||
# 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)
|
||||
|
||||
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 and self.cond_hint_original.size(0) >= self.full_latent_length:
|
||||
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original[self.sub_idxs], x_noisy.shape[3] * 8, x_noisy.shape[2] * 8, 'nearest-exact', "center").to(dtype).to(self.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(self.device)
|
||||
if x_noisy.shape[0] != self.cond_hint.shape[0]:
|
||||
self.cond_hint = broadcast_image_to(self.cond_hint, x_noisy.shape[0], batched_number)
|
||||
# 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.
|
||||
# return normal controlnet stuff
|
||||
return control_prev
|
||||
|
||||
def cleanup_advanced(self):
|
||||
super().cleanup_advanced()
|
||||
# TODO: cleanup patches here
|
||||
|
||||
def copy(self):
|
||||
c = ReferenceAdvanced(self.timestep_keyframes)
|
||||
self.copy_to(c)
|
||||
self.copy_to_advanced(c)
|
||||
return c
|
||||
|
||||
|
||||
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=None):
|
||||
|
||||
@@ -0,0 +1,546 @@
|
||||
from typing import Callable, Union
|
||||
|
||||
import math
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
import comfy.sample
|
||||
import comfy.model_patcher
|
||||
import comfy.utils
|
||||
from comfy.controlnet import ControlBase
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.ldm.modules.attention import BasicTransformerBlock
|
||||
|
||||
from .logger import logger
|
||||
from .utils import (AdvancedControlBase, ControlWeights, TimestepKeyframeGroup, AbstractPreprocWrapper,
|
||||
deepcopy_with_sharing, prepare_mask_batch, broadcast_image_to_full, ddpm_noise_latents, simple_noise_latents)
|
||||
|
||||
|
||||
def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable:
|
||||
def get_refcn(control: ControlBase, order: int=-1):
|
||||
ref_set: set[ReferenceAdvanced] = set()
|
||||
if control is None:
|
||||
return ref_set
|
||||
if type(control) == ReferenceAdvanced:
|
||||
control.order = order
|
||||
order -= 1
|
||||
ref_set.add(control)
|
||||
ref_set.update(get_refcn(control.previous_controlnet, order=order))
|
||||
return ref_set
|
||||
|
||||
def refcn_sample(model: ModelPatcher, *args, **kwargs):
|
||||
# check if positive or negative conds contain ref cn
|
||||
positive = args[-3]
|
||||
negative = args[-2]
|
||||
ref_set = set()
|
||||
if positive is not None:
|
||||
for cond in positive:
|
||||
if "control" in cond[1]:
|
||||
ref_set.update(get_refcn(cond[1]["control"]))
|
||||
if negative is not None:
|
||||
for cond in negative:
|
||||
if "control" in cond[1]:
|
||||
ref_set.update(get_refcn(cond[1]["control"]))
|
||||
# if no ref cn found, do original function immediately
|
||||
if len(ref_set) == 0:
|
||||
return orig_comfy_sample(model, *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])
|
||||
reference_injections.attn_modules = []
|
||||
for i, module in enumerate(attn_modules):
|
||||
injection_holder = InjectionBasicTransformerBlockHolder(block=module, idx=i)
|
||||
injection_holder.attn_weight = float(i) / float(len(attn_modules))
|
||||
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
|
||||
# handle diffusion_model forward injection
|
||||
reference_injections.diffusion_model_orig_forward = model.model.diffusion_model.forward
|
||||
model.model.diffusion_model.forward = factory_forward_inject_UNetModel(reference_injections).__get__(model.model.diffusion_model, type(model.model.diffusion_model))
|
||||
# store ordered ref cns in model's transformer options
|
||||
orig_model_options = model.model_options
|
||||
new_model_options = model.model_options.copy()
|
||||
new_model_options["transformer_options"] = model.model_options["transformer_options"].copy()
|
||||
ref_list: list[ReferenceAdvanced] = list(ref_set)
|
||||
new_model_options["transformer_options"][REF_CONTROL_LIST_ALL] = sorted(ref_list, key=lambda x: x.order)
|
||||
model.model_options = new_model_options
|
||||
# continue with original function
|
||||
return orig_comfy_sample(model, *args, **kwargs)
|
||||
finally:
|
||||
# cleanup injections
|
||||
# first, 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()
|
||||
del module.injection_holder
|
||||
del attn_modules
|
||||
# restore diffusion_model forward function
|
||||
model.model.diffusion_model.forward = reference_injections.diffusion_model_orig_forward.__get__(model.model.diffusion_model, type(model.model.diffusion_model))
|
||||
# restore model_options
|
||||
model.model_options = orig_model_options
|
||||
# cleanup
|
||||
reference_injections.cleanup()
|
||||
return refcn_sample
|
||||
# inject sample functions
|
||||
comfy.sample.sample = refcn_sample_factory(comfy.sample.sample)
|
||||
comfy.sample.sample_custom = refcn_sample_factory(comfy.sample.sample_custom, is_custom=True)
|
||||
|
||||
|
||||
REF_CONTROL_LIST = "ref_control_list"
|
||||
REF_CONTROL_LIST_ALL = "ref_control_list_all"
|
||||
REF_CONTROL_INFO = "ref_control_info"
|
||||
REF_MACHINE_STATE = "ref_machine_state"
|
||||
REF_COND_IDXS = "ref_cond_idxs"
|
||||
REF_UNCOND_IDXS = "ref_uncond_idxs"
|
||||
|
||||
|
||||
class MachineState:
|
||||
WRITE = "write"
|
||||
READ = "read"
|
||||
STYLEALIGN = "stylealign"
|
||||
TEST = "test"
|
||||
|
||||
|
||||
class ReferenceType:
|
||||
ATTN = "reference_attn"
|
||||
ADAIN = "reference_adain"
|
||||
ATTN_ADAIN = "reference_attn+adain"
|
||||
STYLE_ALIGN = "StyleAlign"
|
||||
|
||||
_LIST = [ATTN]
|
||||
_LIST_FULL = [ATTN, ADAIN, ATTN_ADAIN]
|
||||
|
||||
|
||||
class ReferenceOptions:
|
||||
def __init__(self, reference_type: str, style_fidelity: float, ref_weight: float, ref_with_other_cns: bool=False):
|
||||
self.reference_type = reference_type
|
||||
self.original_style_fidelity = style_fidelity
|
||||
self.style_fidelity = style_fidelity
|
||||
self.ref_weight = ref_weight
|
||||
self.ref_with_other_cns = ref_with_other_cns
|
||||
|
||||
def clone(self):
|
||||
return ReferenceOptions(reference_type=self.reference_type, style_fidelity=self.original_style_fidelity, ref_weight=self.ref_weight,
|
||||
ref_with_other_cns=self.ref_with_other_cns)
|
||||
|
||||
|
||||
class ReferencePreprocWrapper(AbstractPreprocWrapper):
|
||||
error_msg = error_msg = "Invalid use of Reference Preprocess output. The output of RGB SparseCtrl preprocessor is NOT a usual image, but a latent pretending to be an image - 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: Tensor):
|
||||
super().__init__(condhint)
|
||||
|
||||
|
||||
class ReferenceAdvanced(ControlBase, AdvancedControlBase):
|
||||
CHANNEL_TO_MULT = {320: 1, 640: 2, 1280: 4}
|
||||
|
||||
def __init__(self, ref_opts: ReferenceOptions, timestep_keyframes: TimestepKeyframeGroup, device=None):
|
||||
super().__init__(device)
|
||||
AdvancedControlBase.__init__(self, super(), timestep_keyframes=timestep_keyframes, weights_default=ControlWeights.controllllite())
|
||||
self.ref_opts = ref_opts
|
||||
self.order = 0
|
||||
self.latent_format = None
|
||||
self.model_sampling_current = None
|
||||
self.should_apply_effective_strength = False
|
||||
self.should_apply_effective_masks = False
|
||||
self.latent_shape = None
|
||||
|
||||
def any_strength_to_apply(self):
|
||||
return self.should_apply_effective_strength or self.should_apply_effective_masks
|
||||
|
||||
def get_effective_strength(self):
|
||||
effective_strength = self.strength
|
||||
if self.current_timestep_keyframe is not None:
|
||||
effective_strength = effective_strength * self.current_timestep_keyframe.strength
|
||||
return effective_strength
|
||||
|
||||
def get_effective_mask_or_float(self, x: Tensor, channels: int, is_mid: bool):
|
||||
if not self.should_apply_effective_masks:
|
||||
return self.get_effective_strength()
|
||||
if is_mid:
|
||||
div = 8
|
||||
else:
|
||||
div = self.CHANNEL_TO_MULT[channels]
|
||||
real_mask = torch.ones([self.latent_shape[0], 1, self.latent_shape[2]//div, self.latent_shape[3]//div]).to(dtype=x.dtype, device=x.device) * self.strength
|
||||
self.apply_advanced_strengths_and_masks(x=real_mask, batched_number=self.batched_number)
|
||||
# mask is now shape [b, 1, h ,w]; need to turn into [b, h*w, 1]
|
||||
b, c, h, w = real_mask.shape
|
||||
real_mask = real_mask.permute(0, 2, 3, 1).reshape(b, h*w, c)
|
||||
return real_mask
|
||||
|
||||
def pre_run_advanced(self, model, percent_to_timestep_function):
|
||||
AdvancedControlBase.pre_run_advanced(self, model, percent_to_timestep_function)
|
||||
if type(self.cond_hint_original) == ReferencePreprocWrapper:
|
||||
self.cond_hint_original = self.cond_hint_original.condhint
|
||||
self.latent_format = model.latent_format # LatentFormat object, used to process_in latent cond_hint
|
||||
self.model_sampling_current = model.model_sampling
|
||||
# SDXL is more sensitive to style_fidelity according to sd-webui-controlnet comments
|
||||
if type(model).__name__ == "SDXL":
|
||||
self.ref_opts.style_fidelity = self.ref_opts.original_style_fidelity ** 3.0
|
||||
else:
|
||||
self.ref_opts.style_fidelity = self.ref_opts.original_style_fidelity
|
||||
|
||||
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int):
|
||||
# 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)
|
||||
|
||||
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 - it is a latent, NOT an image
|
||||
#if self.sub_idxs is not None or self.cond_hint is None or x_noisy.shape[2] != self.cond_hint.shape[2] or x_noisy.shape[3] != 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 and self.cond_hint_original.size(0) >= self.full_latent_length:
|
||||
self.cond_hint = comfy.utils.common_upscale(
|
||||
self.cond_hint_original[self.sub_idxs],
|
||||
x_noisy.shape[3], x_noisy.shape[2], 'nearest-exact', "center").to(dtype).to(self.device)
|
||||
else:
|
||||
self.cond_hint = comfy.utils.common_upscale(
|
||||
self.cond_hint_original,
|
||||
x_noisy.shape[3], x_noisy.shape[2], 'nearest-exact', "center").to(dtype).to(self.device)
|
||||
if x_noisy.shape[0] != self.cond_hint.shape[0]:
|
||||
self.cond_hint = broadcast_image_to_full(self.cond_hint, x_noisy.shape[0], batched_number, except_one=False)
|
||||
# noise cond_hint based on sigma (current step)
|
||||
self.cond_hint = self.latent_format.process_in(self.cond_hint)
|
||||
self.cond_hint = ddpm_noise_latents(self.cond_hint, sigma=t, noise=None)
|
||||
timestep = self.model_sampling_current.timestep(t)
|
||||
self.should_apply_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self.current_timestep_keyframe.strength, 1.0))
|
||||
# prepare mask - use direct_attn, so the mask dims will match source latents (and be smaller)
|
||||
self.prepare_mask_cond_hint(x_noisy=x_noisy, t=t, cond=cond, batched_number=batched_number, direct_attn=True)
|
||||
self.should_apply_effective_masks = self.latent_keyframes is not None or self.mask_cond_hint is not None or self.tk_mask_cond_hint is not None
|
||||
self.latent_shape = list(x_noisy.shape)
|
||||
# done preparing; model patches will take care of everything now.
|
||||
# return normal controlnet stuff
|
||||
return control_prev
|
||||
|
||||
def cleanup_advanced(self):
|
||||
super().cleanup_advanced()
|
||||
del self.latent_format
|
||||
self.latent_format = None
|
||||
del self.model_sampling_current
|
||||
self.model_sampling_current = None
|
||||
self.should_apply_effective_strength = False
|
||||
self.should_apply_effective_masks = False
|
||||
|
||||
def copy(self):
|
||||
c = ReferenceAdvanced(self.ref_opts, self.timestep_keyframes)
|
||||
c.order = self.order
|
||||
self.copy_to(c)
|
||||
self.copy_to_advanced(c)
|
||||
return c
|
||||
|
||||
# avoid deepcopy shenanigans by making deepcopy not do anything to the reference
|
||||
# TODO: do the bookkeeping to do this in a proper way for all Adv-ControlNets
|
||||
def __deepcopy__(self, memo):
|
||||
return self
|
||||
|
||||
|
||||
class BankStylesBasicTransformerBlock:
|
||||
def __init__(self):
|
||||
self.bank = []
|
||||
self.style_cfgs = []
|
||||
self.cn_idx: list[int] = []
|
||||
|
||||
def get_avg_style_fidelity(self):
|
||||
return sum(self.style_cfgs) / float(len(self.style_cfgs))
|
||||
|
||||
def clean(self):
|
||||
del self.bank
|
||||
self.bank = []
|
||||
del self.style_cfgs
|
||||
self.style_cfgs = []
|
||||
del self.cn_idx
|
||||
self.cn_idx = []
|
||||
|
||||
|
||||
class InjectionBasicTransformerBlockHolder:
|
||||
def __init__(self, block: BasicTransformerBlock, idx=None):
|
||||
self.original_forward = block._forward
|
||||
self.idx = idx
|
||||
self.attn_weight = 1.0
|
||||
self.is_middle = False
|
||||
self.bank_styles = BankStylesBasicTransformerBlock()
|
||||
|
||||
def restore(self, block: BasicTransformerBlock):
|
||||
block._forward = self.original_forward
|
||||
|
||||
def clean(self):
|
||||
self.bank_styles.clean()
|
||||
|
||||
|
||||
class ReferenceInjections:
|
||||
def __init__(self, attn_modules: list['RefBasicTransformerBlock']=None):
|
||||
self.attn_modules = attn_modules if attn_modules else []
|
||||
self.diffusion_model_orig_forward: Callable = None
|
||||
|
||||
def clean_module_mem(self):
|
||||
for attn_module in self.attn_modules:
|
||||
try:
|
||||
attn_module.injection_holder.clean()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def cleanup(self):
|
||||
self.clean_module_mem()
|
||||
del self.attn_modules
|
||||
self.attn_modules = []
|
||||
self.diffusion_model_orig_forward = None
|
||||
|
||||
|
||||
def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections):
|
||||
def forward_inject_UNetModel(self, x: Tensor, *args, **kwargs):
|
||||
# get control and transformer_options from kwargs
|
||||
real_args = list(args)
|
||||
real_kwargs = list(kwargs.keys())
|
||||
control = kwargs.get("control", None)
|
||||
transformer_options = kwargs.get("transformer_options", None)
|
||||
# look for ReferenceAttnPatch objects to get ReferenceAdvanced objects
|
||||
ref_controlnets: list[ReferenceAdvanced] = transformer_options[REF_CONTROL_LIST_ALL]
|
||||
# discard any controlnets that should not run
|
||||
ref_controlnets = [x for x in ref_controlnets if x.should_run()]
|
||||
# if nothing related to reference controlnets, do nothing special
|
||||
if len(ref_controlnets) == 0:
|
||||
return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs)
|
||||
try:
|
||||
# assign cond and uncond idxs
|
||||
batched_number = len(transformer_options["cond_or_uncond"])
|
||||
per_batch = x.shape[0] // batched_number
|
||||
indiv_conds = []
|
||||
for cond_type in transformer_options["cond_or_uncond"]:
|
||||
indiv_conds.extend([cond_type] * per_batch)
|
||||
transformer_options[REF_UNCOND_IDXS] = [i for i, x in enumerate(indiv_conds) if x == 1]
|
||||
transformer_options[REF_COND_IDXS] = [i for i, x in enumerate(indiv_conds) if x == 0]
|
||||
# handle running diffusion with ref cond hints
|
||||
for control in ref_controlnets:
|
||||
transformer_options[REF_MACHINE_STATE] = MachineState.WRITE
|
||||
transformer_options[REF_CONTROL_LIST] = [control]
|
||||
# handle masks - apply x to unmasked
|
||||
#strength_mask = torch.ones_like(x, dtype=x.dtype) * control.strength
|
||||
#control.apply_advanced_strengths_and_masks(x=strength_mask, batched_number=batched_number)
|
||||
#real_cond_hint = control.cond_hint * strength_mask + x * (1 - strength_mask)
|
||||
orig_kwargs = kwargs
|
||||
if not control.ref_opts.ref_with_other_cns:
|
||||
kwargs = kwargs.copy()
|
||||
kwargs["control"] = None
|
||||
reference_injections.diffusion_model_orig_forward(control.cond_hint.to(dtype=x.dtype).to(device=x.device), *args, **kwargs)
|
||||
kwargs = orig_kwargs
|
||||
# run diffusion for real now
|
||||
transformer_options[REF_MACHINE_STATE] = MachineState.READ
|
||||
transformer_options[REF_CONTROL_LIST] = ref_controlnets
|
||||
return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs)
|
||||
finally:
|
||||
# make sure banks are cleared no matter what happens - otherwise, RIP VRAM
|
||||
reference_injections.clean_module_mem()
|
||||
|
||||
return forward_inject_UNetModel
|
||||
|
||||
|
||||
# dummy class just to help IDE keep track of injected variables
|
||||
class RefBasicTransformerBlock(BasicTransformerBlock):
|
||||
injection_holder: InjectionBasicTransformerBlockHolder = None
|
||||
|
||||
def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Tensor, context: Tensor=None, transformer_options: dict[str]={}):
|
||||
extra_options = {}
|
||||
block = transformer_options.get("block", None)
|
||||
block_index = transformer_options.get("block_index", 0)
|
||||
transformer_patches = {}
|
||||
transformer_patches_replace = {}
|
||||
|
||||
for k in transformer_options:
|
||||
if k == "patches":
|
||||
transformer_patches = transformer_options[k]
|
||||
elif k == "patches_replace":
|
||||
transformer_patches_replace = transformer_options[k]
|
||||
else:
|
||||
extra_options[k] = transformer_options[k]
|
||||
|
||||
extra_options["n_heads"] = self.n_heads
|
||||
extra_options["dim_head"] = self.d_head
|
||||
|
||||
if self.ff_in:
|
||||
x_skip = x
|
||||
x = self.ff_in(self.norm_in(x))
|
||||
if self.is_res:
|
||||
x += x_skip
|
||||
|
||||
n: Tensor = self.norm1(x)
|
||||
if self.disable_self_attn:
|
||||
context_attn1 = context
|
||||
else:
|
||||
context_attn1 = None
|
||||
value_attn1 = None
|
||||
|
||||
# Reference CN stuff
|
||||
uc_idx_mask = transformer_options.get(REF_UNCOND_IDXS, [])
|
||||
c_idx_mask = transformer_options.get(REF_COND_IDXS, [])
|
||||
# WRITE mode will only have one ReferenceAdvanced, other modes will have all ReferenceAdvanced
|
||||
ref_controlnets: list[ReferenceAdvanced] = transformer_options.get(REF_CONTROL_LIST, None)
|
||||
ref_machine_state: str = transformer_options.get(REF_MACHINE_STATE, None)
|
||||
# if in WRITE mode, save n and style_fidelity
|
||||
if ref_controlnets and ref_machine_state == MachineState.WRITE:
|
||||
if ref_controlnets[0].ref_opts.ref_weight > self.injection_holder.attn_weight:
|
||||
self.injection_holder.bank_styles.bank.append(n.detach().clone())
|
||||
self.injection_holder.bank_styles.style_cfgs.append(ref_controlnets[0].ref_opts.style_fidelity)
|
||||
self.injection_holder.bank_styles.cn_idx.append(ref_controlnets[0].order)
|
||||
|
||||
if "attn1_patch" in transformer_patches:
|
||||
patch = transformer_patches["attn1_patch"]
|
||||
if context_attn1 is None:
|
||||
context_attn1 = n
|
||||
value_attn1 = context_attn1
|
||||
for p in patch:
|
||||
n, context_attn1, value_attn1 = p(n, context_attn1, value_attn1, extra_options)
|
||||
|
||||
if block is not None:
|
||||
transformer_block = (block[0], block[1], block_index)
|
||||
else:
|
||||
transformer_block = None
|
||||
attn1_replace_patch = transformer_patches_replace.get("attn1", {})
|
||||
block_attn1 = transformer_block
|
||||
if block_attn1 not in attn1_replace_patch:
|
||||
block_attn1 = block
|
||||
|
||||
if block_attn1 in attn1_replace_patch:
|
||||
if context_attn1 is None:
|
||||
context_attn1 = n
|
||||
value_attn1 = n
|
||||
n = self.attn1.to_q(n)
|
||||
# Reference CN READ - use attn1_replace_patch appropriately
|
||||
if ref_machine_state == MachineState.READ and len(self.injection_holder.bank_styles.bank) > 0:
|
||||
bank_styles = self.injection_holder.bank_styles
|
||||
style_fidelity = bank_styles.get_avg_style_fidelity()
|
||||
real_bank = bank_styles.bank.copy()
|
||||
for idx, order in enumerate(bank_styles.cn_idx):
|
||||
if ref_controlnets[idx].any_strength_to_apply():
|
||||
effective_strength = ref_controlnets[idx].get_effective_mask_or_float(x=n, channels=n.shape[2], is_mid=self.injection_holder.is_middle)
|
||||
real_bank[idx] = real_bank[idx] * effective_strength + context_attn1 * (1-effective_strength)
|
||||
n_uc = self.attn1.to_out(attn1_replace_patch[block_attn1](
|
||||
n,
|
||||
self.attn1.to_k(torch.cat([context_attn1] + real_bank, dim=1)),
|
||||
self.attn1.to_v(torch.cat([value_attn1] + real_bank, dim=1)),
|
||||
extra_options))
|
||||
n_c = n_uc.clone()
|
||||
if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0):
|
||||
n_c[uc_idx_mask] = self.attn1.to_out(attn1_replace_patch[block_attn1](
|
||||
n[uc_idx_mask],
|
||||
self.attn1.to_k(context_attn1[uc_idx_mask]),
|
||||
self.attn1.to_v(value_attn1[uc_idx_mask]),
|
||||
extra_options))
|
||||
n = style_fidelity * n_c + (1.0-style_fidelity) * n_uc
|
||||
bank_styles.clean()
|
||||
else:
|
||||
context_attn1 = self.attn1.to_k(context_attn1)
|
||||
value_attn1 = self.attn1.to_v(value_attn1)
|
||||
n = attn1_replace_patch[block_attn1](n, context_attn1, value_attn1, extra_options)
|
||||
n = self.attn1.to_out(n)
|
||||
else:
|
||||
# Reference CN READ - no attn1_replace_patch
|
||||
if ref_machine_state == MachineState.READ and len(self.injection_holder.bank_styles.bank) > 0:
|
||||
if context_attn1 is None:
|
||||
context_attn1 = n
|
||||
bank_styles = self.injection_holder.bank_styles
|
||||
style_fidelity = bank_styles.get_avg_style_fidelity()
|
||||
real_bank = bank_styles.bank.copy()
|
||||
for idx, order in enumerate(bank_styles.cn_idx):
|
||||
if ref_controlnets[idx].any_strength_to_apply():
|
||||
effective_strength = ref_controlnets[idx].get_effective_mask_or_float(x=n, channels=n.shape[2], is_mid=self.injection_holder.is_middle)
|
||||
real_bank[idx] = real_bank[idx] * effective_strength + context_attn1 * (1-effective_strength)
|
||||
n_uc: Tensor = self.attn1(
|
||||
n,
|
||||
context=torch.cat([context_attn1] + real_bank, dim=1),
|
||||
value=torch.cat([value_attn1] + real_bank, dim=1) if value_attn1 is not None else value_attn1)
|
||||
n_c = n_uc.clone()
|
||||
if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0):
|
||||
n_c[uc_idx_mask] = self.attn1(
|
||||
n[uc_idx_mask],
|
||||
context=context_attn1[uc_idx_mask],
|
||||
value=value_attn1[uc_idx_mask] if value_attn1 is not None else value_attn1)
|
||||
n = style_fidelity * n_c + (1.0-style_fidelity) * n_uc
|
||||
bank_styles.clean()
|
||||
else:
|
||||
n = self.attn1(n, context=context_attn1, value=value_attn1)
|
||||
|
||||
if "attn1_output_patch" in transformer_patches:
|
||||
patch = transformer_patches["attn1_output_patch"]
|
||||
for p in patch:
|
||||
n = p(n, extra_options)
|
||||
|
||||
x += n
|
||||
if "middle_patch" in transformer_patches:
|
||||
patch = transformer_patches["middle_patch"]
|
||||
for p in patch:
|
||||
x = p(x, extra_options)
|
||||
|
||||
if self.attn2 is not None:
|
||||
n = self.norm2(x)
|
||||
if self.switch_temporal_ca_to_sa:
|
||||
context_attn2 = n
|
||||
else:
|
||||
context_attn2 = context
|
||||
value_attn2 = None
|
||||
if "attn2_patch" in transformer_patches:
|
||||
patch = transformer_patches["attn2_patch"]
|
||||
value_attn2 = context_attn2
|
||||
for p in patch:
|
||||
n, context_attn2, value_attn2 = p(n, context_attn2, value_attn2, extra_options)
|
||||
|
||||
attn2_replace_patch = transformer_patches_replace.get("attn2", {})
|
||||
block_attn2 = transformer_block
|
||||
if block_attn2 not in attn2_replace_patch:
|
||||
block_attn2 = block
|
||||
|
||||
if block_attn2 in attn2_replace_patch:
|
||||
if value_attn2 is None:
|
||||
value_attn2 = context_attn2
|
||||
n = self.attn2.to_q(n)
|
||||
context_attn2 = self.attn2.to_k(context_attn2)
|
||||
value_attn2 = self.attn2.to_v(value_attn2)
|
||||
n = attn2_replace_patch[block_attn2](n, context_attn2, value_attn2, extra_options)
|
||||
n = self.attn2.to_out(n)
|
||||
else:
|
||||
n = self.attn2(n, context=context_attn2, value=value_attn2)
|
||||
|
||||
if "attn2_output_patch" in transformer_patches:
|
||||
patch = transformer_patches["attn2_output_patch"]
|
||||
for p in patch:
|
||||
n = p(n, extra_options)
|
||||
|
||||
x += n
|
||||
if self.is_res:
|
||||
x_skip = x
|
||||
x = self.ff(self.norm3(x))
|
||||
if self.is_res:
|
||||
x += x_skip
|
||||
|
||||
return x
|
||||
|
||||
|
||||
# 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
|
||||
|
||||
@@ -11,6 +11,7 @@ from .nodes_weight import (DefaultWeights, ScaledSoftMaskedUniversalWeights, Sca
|
||||
SoftT2IAdapterWeights, CustomT2IAdapterWeights)
|
||||
from .nodes_latent_keyframe import LatentKeyframeGroupNode, LatentKeyframeInterpolationNode, LatentKeyframeBatchedGroupNode, LatentKeyframeNode
|
||||
from .nodes_sparsectrl import SparseCtrlMergedLoaderAdvanced, SparseCtrlLoaderAdvanced, SparseIndexMethodNode, SparseSpreadMethodNode, RgbSparseCtrlPreprocessor
|
||||
from .nodes_reference import ReferenceControlNetNode, ReferencePreprocessorNode
|
||||
from .nodes_loosecontrol import ControlNetLoaderWithLoraAdvanced
|
||||
from .nodes_deprecated import LoadImagesFromDirectory
|
||||
from .logger import logger
|
||||
@@ -233,6 +234,9 @@ NODE_CLASS_MAPPINGS = {
|
||||
"ACN_SparseCtrlMergedLoaderAdvanced": SparseCtrlMergedLoaderAdvanced,
|
||||
"ACN_SparseCtrlIndexMethodNode": SparseIndexMethodNode,
|
||||
"ACN_SparseCtrlSpreadMethodNode": SparseSpreadMethodNode,
|
||||
# Reference
|
||||
"ACN_ReferencePreprocessor": ReferencePreprocessorNode,
|
||||
"ACN_ReferenceControlNet": ReferenceControlNetNode,
|
||||
# LOOSEControl
|
||||
#"ACN_ControlNetLoaderWithLoraAdvanced": ControlNetLoaderWithLoraAdvanced,
|
||||
# Deprecated
|
||||
@@ -265,8 +269,11 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"ACN_SparseCtrlMergedLoaderAdvanced": "Load Merged SparseCtrl Model 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlIndexMethodNode": "SparseCtrl Index Method 🛂🅐🅒🅝",
|
||||
"ACN_SparseCtrlSpreadMethodNode": "SparseCtrl Spread Method 🛂🅐🅒🅝",
|
||||
# Reference
|
||||
"ACN_ReferencePreprocessor": "Reference Preproccessor 🛂🅐🅒🅝",
|
||||
"ACN_ReferenceControlNet": "Reference ControlNet 🛂🅐🅒🅝",
|
||||
# LOOSEControl
|
||||
#"ACN_ControlNetLoaderWithLoraAdvanced": "Load Adv. ControlNet Model w/ LoRA 🛂🅐🅒🅝",
|
||||
# Deprecated
|
||||
"LoadImagesFromDirectory": "Load Images [DEPRECATED] 🛂🅐🅒🅝",
|
||||
"LoadImagesFromDirectory": "🚫Load Images [DEPRECATED] 🛂🅐🅒🅝",
|
||||
}
|
||||
|
||||
@@ -0,0 +1,62 @@
|
||||
from torch import Tensor
|
||||
|
||||
from nodes import VAEEncode
|
||||
import comfy.utils
|
||||
from comfy.sd import VAE
|
||||
|
||||
from .control_reference import ReferenceAdvanced, ReferenceOptions, ReferenceType, ReferencePreprocWrapper
|
||||
|
||||
|
||||
# node for ReferenceCN
|
||||
class ReferenceControlNetNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"reference_type": (ReferenceType._LIST,),
|
||||
"style_fidelity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"ref_weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET", )
|
||||
FUNCTION = "load_controlnet"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/Reference"
|
||||
|
||||
def load_controlnet(self, reference_type: str, style_fidelity: float, ref_weight: float, ref_with_other_cns: bool=False):
|
||||
ref_opts = ReferenceOptions(reference_type=reference_type, style_fidelity=style_fidelity, ref_weight=ref_weight,
|
||||
ref_with_other_cns=ref_with_other_cns)
|
||||
controlnet = ReferenceAdvanced(ref_opts=ref_opts, timestep_keyframes=None)
|
||||
return (controlnet,)
|
||||
|
||||
|
||||
class ReferencePreprocessorNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE", ),
|
||||
"vae": ("VAE", ),
|
||||
"latent_size": ("LATENT", ),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("proc_IMAGE",)
|
||||
FUNCTION = "preprocess_images"
|
||||
|
||||
CATEGORY = "Adv-ControlNet 🛂🅐🅒🅝/Reference/preprocess"
|
||||
|
||||
def preprocess_images(self, vae: VAE, image: Tensor, latent_size: Tensor):
|
||||
# first, resize image to match latents
|
||||
image = image.movedim(-1,1)
|
||||
image = comfy.utils.common_upscale(image, latent_size["samples"].shape[3] * 8, latent_size["samples"].shape[2] * 8, 'nearest-exact', "center")
|
||||
image = image.movedim(1,-1)
|
||||
# then, vae encode
|
||||
try:
|
||||
image = vae.vae_encode_crop_pixels(image)
|
||||
except Exception:
|
||||
image = VAEEncode.vae_encode_crop_pixels(image)
|
||||
encoded = vae.encode(image[:,:,:,:3])
|
||||
return (ReferencePreprocWrapper(condhint=encoded),)
|
||||
|
||||
+90
-9
@@ -3,6 +3,7 @@ from typing import Callable, Union
|
||||
import torch
|
||||
from torch import Tensor
|
||||
import torch.nn.functional as F
|
||||
import math
|
||||
|
||||
import comfy.ops
|
||||
import comfy.utils
|
||||
@@ -11,8 +12,8 @@ from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .logger import logger
|
||||
|
||||
BIGMIN = -(2**63-1)
|
||||
BIGMAX = (2**63-1)
|
||||
BIGMIN = -(2**53-1)
|
||||
BIGMAX = (2**53-1)
|
||||
|
||||
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):
|
||||
@@ -232,6 +233,38 @@ class TimestepKeyframeGroup:
|
||||
return group
|
||||
|
||||
|
||||
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: Tensor):
|
||||
self.condhint = condhint
|
||||
|
||||
def movedim(self, *args, **kwargs):
|
||||
return self
|
||||
|
||||
def __getattr__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __setattr__(self, name, value):
|
||||
if name != "condhint":
|
||||
raise AttributeError(self.error_msg)
|
||||
super().__setattr__(name, value)
|
||||
|
||||
def __iter__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __next__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __len__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __getitem__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
def __setitem__(self, *args, **kwargs):
|
||||
raise AttributeError(self.error_msg)
|
||||
|
||||
|
||||
# depending on model, AnimateDiff may inject into GroupNorm, so make sure GroupNorm will be clean
|
||||
class disable_weight_init_clean_groupnorm(comfy.ops.disable_weight_init):
|
||||
class GroupNorm(comfy.ops.disable_weight_init.GroupNorm):
|
||||
@@ -264,6 +297,46 @@ def linear_conversion(x, x_min=0.0, x_max=1.0, new_min=0.0, new_max=1.0):
|
||||
return (((x - x_min)/(x_max - x_min)) * (new_max - new_min)) + new_min
|
||||
|
||||
|
||||
def broadcast_image_to_full(tensor, target_batch_size, batched_number, except_one=True):
|
||||
current_batch_size = tensor.shape[0]
|
||||
#print(current_batch_size, target_batch_size)
|
||||
if except_one and current_batch_size == 1:
|
||||
return tensor
|
||||
|
||||
per_batch = target_batch_size // batched_number
|
||||
tensor = tensor[:per_batch]
|
||||
|
||||
if per_batch > tensor.shape[0]:
|
||||
tensor = torch.cat([tensor] * (per_batch // tensor.shape[0]) + [tensor[:(per_batch % tensor.shape[0])]], dim=0)
|
||||
|
||||
current_batch_size = tensor.shape[0]
|
||||
if current_batch_size == target_batch_size:
|
||||
return tensor
|
||||
else:
|
||||
return torch.cat([tensor] * batched_number, dim=0)
|
||||
|
||||
|
||||
def ddpm_noise_latents(latents: Tensor, sigma: Tensor, noise: Tensor=None):
|
||||
sigma = sigma.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
|
||||
alpha_cumprod = 1 / ((sigma * sigma) + 1)
|
||||
sqrt_alpha_prod = alpha_cumprod ** 0.5
|
||||
sqrt_one_minus_alpha_prod = (1. - alpha_cumprod) ** 0.5
|
||||
if noise is None:
|
||||
# generator = torch.Generator(device="cuda")
|
||||
# generator.manual_seed(0)
|
||||
# generator = torch.Generator()
|
||||
# generator.manual_seed(0)
|
||||
# noise = torch.randn(latents.size(), generator=generator).to(latents.device)
|
||||
noise = torch.randn_like(latents).to(latents.device)
|
||||
return sqrt_alpha_prod * latents + sqrt_one_minus_alpha_prod * noise
|
||||
|
||||
|
||||
def simple_noise_latents(latents: Tensor, sigma: float, noise: Tensor=None):
|
||||
if noise is None:
|
||||
noise = torch.rand_like(latents)
|
||||
return latents + noise * sigma
|
||||
|
||||
|
||||
# from https://stackoverflow.com/a/24621200
|
||||
def deepcopy_with_sharing(obj, shared_attribute_names, memo=None):
|
||||
'''
|
||||
@@ -458,6 +531,14 @@ class AdvancedControlBase:
|
||||
def disarm(self):
|
||||
self.disarmed = True
|
||||
|
||||
def should_run(self):
|
||||
if math.isclose(self.strength, 0.0) or math.isclose(self.current_timestep_keyframe.strength, 0.0):
|
||||
return False
|
||||
if self.timestep_range is not None:
|
||||
if self.t > self.timestep_range[0] or self.t < self.timestep_range[1]:
|
||||
return False
|
||||
return True
|
||||
|
||||
def get_control_inject(self, x_noisy, t, cond, batched_number):
|
||||
# prepare timestep and everything related
|
||||
self.prepare_current_timestep(t=t, batched_number=batched_number)
|
||||
@@ -601,12 +682,12 @@ class AdvancedControlBase:
|
||||
o[i] += prev_val
|
||||
return out
|
||||
|
||||
def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None):
|
||||
self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype)
|
||||
self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype)
|
||||
def prepare_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False):
|
||||
self._prepare_mask("mask_cond_hint", self.mask_cond_hint_original, x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn)
|
||||
self.prepare_tk_mask_cond_hint(x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn)
|
||||
|
||||
def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None):
|
||||
return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype)
|
||||
def prepare_tk_mask_cond_hint(self, x_noisy: Tensor, t, cond, batched_number, dtype=None, direct_attn=False):
|
||||
return self._prepare_mask("tk_mask_cond_hint", self.current_timestep_keyframe.mask_hint_orig, x_noisy, t, cond, batched_number, dtype, direct_attn=direct_attn)
|
||||
|
||||
def prepare_weight_mask_cond_hint(self, x_noisy: Tensor, batched_number, dtype=None):
|
||||
return self._prepare_mask("weight_mask_cond_hint", self.weights.weight_mask, x_noisy, t=None, cond=None, batched_number=batched_number, dtype=dtype, direct_attn=True)
|
||||
@@ -615,12 +696,12 @@ class AdvancedControlBase:
|
||||
# make mask appropriate dimensions, if present
|
||||
if orig_mask is not None:
|
||||
out_mask = getattr(self, attr_name)
|
||||
if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * 8 != out_mask.shape[1] or x_noisy.shape[3] * 8 != out_mask.shape[2]:
|
||||
multiplier = 1 if direct_attn else 8
|
||||
if self.sub_idxs is not None or out_mask is None or x_noisy.shape[2] * multiplier != out_mask.shape[1] or x_noisy.shape[3] * multiplier != out_mask.shape[2]:
|
||||
self._reset_attr(attr_name)
|
||||
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
|
||||
multiplier = 1 if direct_attn else 8
|
||||
out_mask = prepare_mask_batch(orig_mask, x_noisy.shape, multiplier=multiplier)
|
||||
actual_latent_length = x_noisy.shape[0] // batched_number
|
||||
out_mask = comfy.utils.repeat_to_batch_size(out_mask, actual_latent_length if self.sub_idxs is None else self.full_latent_length)
|
||||
|
||||
Reference in New Issue
Block a user