Merge PR #93 from Kosinkadink/develop: reference_adain and reference_adain+attn support

Added reference_adain and reference_adain+attn support
This commit is contained in:
Jedrzej Kosinski
2024-04-02 01:15:35 -05:00
committed by GitHub
4 changed files with 355 additions and 68 deletions
+320 -42
View File
@@ -10,10 +10,11 @@ import comfy.utils
from comfy.controlnet import ControlBase
from comfy.model_patcher import ModelPatcher
from comfy.ldm.modules.attention import BasicTransformerBlock
from comfy.ldm.modules.diffusionmodules import openaimodel
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)
deepcopy_with_sharing, prepare_mask_batch, broadcast_image_to_full)
def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable:
@@ -49,6 +50,7 @@ def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callab
# 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] = []
@@ -57,7 +59,6 @@ def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callab
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))
@@ -70,6 +71,37 @@ def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callab
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
# 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))
@@ -84,13 +116,20 @@ def refcn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callab
return orig_comfy_sample(model, *args, **kwargs)
finally:
# cleanup injections
# first, restore attn modules
# 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 gn modules
gn_modules: list[RefTimestepEmbedSequential] = reference_injections.gn_modules
for module in gn_modules:
module.injection_holder.restore(module)
module.injection_holder.clean()
del module.injection_holder
del gn_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
@@ -103,10 +142,12 @@ 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_ATTN_CONTROL_LIST = "ref_attn_control_list"
REF_ADAIN_CONTROL_LIST = "ref_adain_control_list"
REF_CONTROL_LIST_ALL = "ref_control_list_all"
REF_CONTROL_INFO = "ref_control_info"
REF_MACHINE_STATE = "ref_machine_state"
REF_ATTN_MACHINE_STATE = "ref_attn_machine_state"
REF_ADAIN_MACHINE_STATE = "ref_adain_machine_state"
REF_COND_IDXS = "ref_cond_idxs"
REF_UNCOND_IDXS = "ref_uncond_idxs"
@@ -115,7 +156,7 @@ class MachineState:
WRITE = "write"
READ = "read"
STYLEALIGN = "stylealign"
TEST = "test"
OFF = "off"
class ReferenceType:
@@ -124,22 +165,54 @@ class ReferenceType:
ATTN_ADAIN = "reference_attn+adain"
STYLE_ALIGN = "StyleAlign"
_LIST = [ATTN]
_LIST_FULL = [ATTN, ADAIN, ATTN_ADAIN]
_LIST = [ATTN, ADAIN, ATTN_ADAIN]
_LIST_ATTN = [ATTN, ATTN_ADAIN]
_LIST_ADAIN = [ADAIN, ATTN_ADAIN]
@classmethod
def is_attn(cls, ref_type: str):
return ref_type in cls._LIST_ATTN
@classmethod
def is_adain(cls, ref_type: str):
return ref_type in cls._LIST_ADAIN
class ReferenceOptions:
def __init__(self, reference_type: str, style_fidelity: float, ref_weight: float, ref_with_other_cns: bool=False):
def __init__(self, reference_type: str,
attn_style_fidelity: float, adain_style_fidelity: float,
attn_ref_weight: float, adain_ref_weight: float,
attn_strength: float=1.0, adain_strength: float=1.0,
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
# attn
self.original_attn_style_fidelity = attn_style_fidelity
self.attn_style_fidelity = attn_style_fidelity
self.attn_ref_weight = attn_ref_weight
self.attn_strength = attn_strength
# adain
self.original_adain_style_fidelity = adain_style_fidelity
self.adain_style_fidelity = adain_style_fidelity
self.adain_ref_weight = adain_ref_weight
self.adain_strength = adain_strength
# other
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,
return ReferenceOptions(reference_type=self.reference_type,
attn_style_fidelity=self.original_attn_style_fidelity, adain_style_fidelity=self.original_adain_style_fidelity,
attn_ref_weight=self.attn_ref_weight, adain_ref_weight=self.adain_ref_weight,
attn_strength=self.attn_strength, adain_strength=self.adain_strength,
ref_with_other_cns=self.ref_with_other_cns)
@staticmethod
def create_combo(reference_type: str, style_fidelity: float, ref_weight: float, ref_with_other_cns: bool=False):
return ReferenceOptions(reference_type=reference_type,
attn_style_fidelity=style_fidelity, adain_style_fidelity=style_fidelity,
attn_ref_weight=ref_weight, adain_ref_weight=ref_weight,
ref_with_other_cns=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."
@@ -157,33 +230,59 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
self.order = 0
self.latent_format = None
self.model_sampling_current = None
self.should_apply_effective_strength = False
self.should_apply_attn_effective_strength = False
self.should_apply_adain_effective_strength = False
self.should_apply_effective_masks = False
self.latent_shape = None
def any_attn_strength_to_apply(self):
return self.should_apply_attn_effective_strength or self.should_apply_effective_masks
def any_strength_to_apply(self):
return self.should_apply_effective_strength or self.should_apply_effective_masks
def any_adain_strength_to_apply(self):
return self.should_apply_adain_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):
def get_effective_attn_mask_or_float(self, x: Tensor, channels: int, is_mid: bool):
if not self.should_apply_effective_masks:
return self.get_effective_strength()
return self.get_effective_strength() * self.ref_opts.attn_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
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.ref_opts.attn_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 get_effective_adain_mask_or_float(self, x: Tensor):
if not self.should_apply_effective_masks:
return self.get_effective_strength() * self.ref_opts.adain_strength
b, c, h, w = x.shape
real_mask = torch.ones([b, 1, h, w]).to(dtype=x.dtype, device=x.device) * self.strength * self.ref_opts.adain_strength
self.apply_advanced_strengths_and_masks(x=real_mask, batched_number=self.batched_number)
return real_mask
def should_run(self):
running = super().should_run()
if not running:
return running
attn_run = False
adain_run = False
if ReferenceType.is_attn(self.ref_opts.reference_type):
# attn will run as long as neither weight or strength is zero
attn_run = not (math.isclose(self.ref_opts.attn_ref_weight, 0.0) or math.isclose(self.ref_opts.attn_strength, 0.0))
if ReferenceType.is_adain(self.ref_opts.reference_type):
# adain will run as long as neither weight or strength is zero
adain_run = not (math.isclose(self.ref_opts.adain_ref_weight, 0.0) or math.isclose(self.ref_opts.adain_strength, 0.0))
return attn_run or adain_run
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:
@@ -192,9 +291,11 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
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
self.ref_opts.attn_style_fidelity = self.ref_opts.original_attn_style_fidelity ** 3.0
self.ref_opts.adain_style_fidelity = self.ref_opts.original_adain_style_fidelity ** 3.0
else:
self.ref_opts.style_fidelity = self.ref_opts.original_style_fidelity
self.ref_opts.attn_style_fidelity = self.ref_opts.original_attn_style_fidelity
self.ref_opts.adain_style_fidelity = self.ref_opts.original_adain_style_fidelity
def get_control_advanced(self, x_noisy: Tensor, t, cond, batched_number: int):
# normal ControlNet stuff
@@ -225,9 +326,10 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
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)
self.cond_hint = ref_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))
self.should_apply_attn_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self.current_timestep_keyframe.strength, 1.0) and math.isclose(self.ref_opts.attn_strength, 1.0))
self.should_apply_adain_effective_strength = not (math.isclose(self.strength, 1.0) and math.isclose(self.current_timestep_keyframe.strength, 1.0) and math.isclose(self.ref_opts.adain_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
@@ -242,7 +344,8 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
self.latent_format = None
del self.model_sampling_current
self.model_sampling_current = None
self.should_apply_effective_strength = False
self.should_apply_attn_effective_strength = False
self.should_apply_adain_effective_strength = False
self.should_apply_effective_masks = False
def copy(self):
@@ -258,6 +361,28 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
return self
def ref_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)
# noise = torch.empty_like(latents).normal_(generator=generator)
# 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
class BankStylesBasicTransformerBlock:
def __init__(self):
self.bank = []
@@ -276,6 +401,33 @@ class BankStylesBasicTransformerBlock:
self.cn_idx = []
class BankStylesTimestepEmbedSequential:
def __init__(self):
self.var_bank = []
self.mean_bank = []
self.style_cfgs = []
self.cn_idx: list[int] = []
def get_avg_var_bank(self):
return sum(self.var_bank) / float(len(self.var_bank))
def get_avg_mean_bank(self):
return sum(self.mean_bank) / float(len(self.mean_bank))
def get_avg_style_fidelity(self):
return sum(self.style_cfgs) / float(len(self.style_cfgs))
def clean(self):
del self.mean_bank
self.mean_bank = []
del self.var_bank
self.var_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
@@ -291,9 +443,27 @@ class InjectionBasicTransformerBlockHolder:
self.bank_styles.clean()
class InjectionTimestepEmbedSequentialHolder:
def __init__(self, block: openaimodel.TimestepEmbedSequential, idx=None, is_middle=False, is_input=False, is_output=False):
self.original_forward = block.forward
self.idx = idx
self.gn_weight = 1.0
self.is_middle = is_middle
self.is_input = is_input
self.is_output = is_output
self.bank_styles = BankStylesTimestepEmbedSequential()
def restore(self, block: openaimodel.TimestepEmbedSequential):
block.forward = self.original_forward
def clean(self):
self.bank_styles.clean()
class ReferenceInjections:
def __init__(self, attn_modules: list['RefBasicTransformerBlock']=None):
def __init__(self, attn_modules: list['RefBasicTransformerBlock']=None, gn_modules: list['RefTimestepEmbedSequential']=None):
self.attn_modules = attn_modules if attn_modules else []
self.gn_modules = gn_modules if gn_modules else []
self.diffusion_model_orig_forward: Callable = None
def clean_module_mem(self):
@@ -302,11 +472,18 @@ class ReferenceInjections:
attn_module.injection_holder.clean()
except Exception:
pass
for gn_module in self.gn_modules:
try:
gn_module.injection_holder.clean()
except Exception:
pass
def cleanup(self):
self.clean_module_mem()
del self.attn_modules
self.attn_modules = []
del self.gn_modules
self.gn_modules = []
self.diffusion_model_orig_forward = None
@@ -333,14 +510,31 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections):
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]
# check which controlnets do which thing
attn_controlnets = []
adain_controlnets = []
for control in ref_controlnets:
if ReferenceType.is_attn(control.ref_opts.reference_type):
attn_controlnets.append(control)
if ReferenceType.is_adain(control.ref_opts.reference_type):
adain_controlnets.append(control)
if len(adain_controlnets) > 0:
# ComfyUI uses forward_timestep_embed with the TimestepEmbedSequential passed into it
orig_forward_timestep_embed = openaimodel.forward_timestep_embed
openaimodel.forward_timestep_embed = forward_timestep_embed_ref_inject_factory(orig_forward_timestep_embed)
# 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)
if ReferenceType.is_attn(control.ref_opts.reference_type):
transformer_options[REF_ATTN_MACHINE_STATE] = MachineState.WRITE
else:
transformer_options[REF_ATTN_MACHINE_STATE] = MachineState.OFF
if ReferenceType.is_adain(control.ref_opts.reference_type):
transformer_options[REF_ADAIN_MACHINE_STATE] = MachineState.WRITE
else:
transformer_options[REF_ADAIN_MACHINE_STATE] = MachineState.OFF
transformer_options[REF_ATTN_CONTROL_LIST] = [control]
transformer_options[REF_ADAIN_CONTROL_LIST] = [control]
orig_kwargs = kwargs
if not control.ref_opts.ref_with_other_cns:
kwargs = kwargs.copy()
@@ -348,12 +542,16 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections):
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
transformer_options[REF_ATTN_MACHINE_STATE] = MachineState.READ
transformer_options[REF_ADAIN_MACHINE_STATE] = MachineState.READ
transformer_options[REF_ATTN_CONTROL_LIST] = attn_controlnets
transformer_options[REF_ADAIN_CONTROL_LIST] = adain_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()
if len(adain_controlnets) > 0:
openaimodel.forward_timestep_embed = orig_forward_timestep_embed
return forward_inject_UNetModel
@@ -397,13 +595,13 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten
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)
ref_controlnets: list[ReferenceAdvanced] = transformer_options.get(REF_ATTN_CONTROL_LIST, None)
ref_machine_state: str = transformer_options.get(REF_ATTN_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:
if ref_controlnets[0].ref_opts.attn_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.style_cfgs.append(ref_controlnets[0].ref_opts.attn_style_fidelity)
self.injection_holder.bank_styles.cn_idx.append(ref_controlnets[0].order)
if "attn1_patch" in transformer_patches:
@@ -433,9 +631,16 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten
bank_styles = self.injection_holder.bank_styles
style_fidelity = bank_styles.get_avg_style_fidelity()
real_bank = bank_styles.bank.copy()
cn_idx = 0
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)
# make sure matching ref cn is selected
for i in range(cn_idx, len(ref_controlnets)):
if ref_controlnets[i].order == order:
cn_idx = i
break
assert order == ref_controlnets[cn_idx].order
if ref_controlnets[cn_idx].any_attn_strength_to_apply():
effective_strength = ref_controlnets[cn_idx].get_effective_attn_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,
@@ -464,9 +669,16 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten
bank_styles = self.injection_holder.bank_styles
style_fidelity = bank_styles.get_avg_style_fidelity()
real_bank = bank_styles.bank.copy()
cn_idx = 0
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)
# make sure matching ref cn is selected
for i in range(cn_idx, len(ref_controlnets)):
if ref_controlnets[i].order == order:
cn_idx = i
break
assert order == ref_controlnets[cn_idx].order
if ref_controlnets[cn_idx].any_attn_strength_to_apply():
effective_strength = ref_controlnets[cn_idx].get_effective_attn_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,
@@ -538,6 +750,72 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten
return x
class RefTimestepEmbedSequential(openaimodel.TimestepEmbedSequential):
injection_holder: InjectionTimestepEmbedSequentialHolder = None
def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory: Callable):
def forward_timestep_embed_ref_inject(*args, **kwargs):
ts: RefTimestepEmbedSequential = args[0]
if not hasattr(ts, "injection_holder"):
return orig_timestep_embed_inject_factory(*args, **kwargs)
eps = 1e-6
x: Tensor = orig_timestep_embed_inject_factory(*args, **kwargs)
y: Tensor = None
transformer_options: dict[str] = args[4]
# 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_ADAIN_CONTROL_LIST, None)
ref_machine_state: str = transformer_options.get(REF_ADAIN_MACHINE_STATE, None)
# if in WRITE mode, save var, mean, and style_cfg
if ref_machine_state == MachineState.WRITE:
if ref_controlnets[0].ref_opts.adain_ref_weight > ts.injection_holder.gn_weight:
var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0)
ts.injection_holder.bank_styles.var_bank.append(var)
ts.injection_holder.bank_styles.mean_bank.append(mean)
ts.injection_holder.bank_styles.style_cfgs.append(ref_controlnets[0].ref_opts.adain_style_fidelity)
ts.injection_holder.bank_styles.cn_idx.append(ref_controlnets[0].order)
# if in READ mode, do math with saved var, mean, and style_cfg
if ref_machine_state == MachineState.READ:
if len(ts.injection_holder.bank_styles.var_bank) > 0:
bank_styles = ts.injection_holder.bank_styles
var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0)
std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5
y_uc = torch.zeros_like(x)
cn_idx = 0
for idx, order in enumerate(bank_styles.cn_idx):
# make sure matching ref cn is selected
for i in range(cn_idx, len(ref_controlnets)):
if ref_controlnets[i].order == order:
cn_idx = i
break
assert order == ref_controlnets[cn_idx].order
style_fidelity = bank_styles.style_cfgs[idx]
var_acc = bank_styles.var_bank[idx]
mean_acc = bank_styles.mean_bank[idx]
std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5
sub_y_uc = (((x - mean) / std) * std_acc) + mean_acc
if ref_controlnets[cn_idx].any_adain_strength_to_apply():
effective_strength = ref_controlnets[cn_idx].get_effective_adain_mask_or_float(x=x)
sub_y_uc = sub_y_uc * effective_strength + x * (1-effective_strength)
y_uc += sub_y_uc
# get average, if more than one
if len(bank_styles.cn_idx) > 1:
y_uc /= len(bank_styles.cn_idx)
y_c = y_uc.clone()
if len(uc_idx_mask) > 0 and not math.isclose(style_fidelity, 0.0):
y_c[uc_idx_mask] = x.to(y_c.dtype)[uc_idx_mask]
y = style_fidelity * y_c + (1.0 - style_fidelity) * y_uc
ts.injection_holder.bank_styles.clean()
if y is None:
y = x
return y.to(x.dtype)
return forward_timestep_embed_ref_inject
# DFS Search for Torch.nn.Module, Written by Lvmin
def torch_dfs(model: torch.nn.Module):
result = [model]
+4 -2
View File
@@ -11,7 +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_reference import ReferenceControlNetNode, ReferenceControlFinetune, ReferencePreprocessorNode
from .nodes_loosecontrol import ControlNetLoaderWithLoraAdvanced
from .nodes_deprecated import LoadImagesFromDirectory
from .logger import logger
@@ -237,6 +237,7 @@ NODE_CLASS_MAPPINGS = {
# Reference
"ACN_ReferencePreprocessor": ReferencePreprocessorNode,
"ACN_ReferenceControlNet": ReferenceControlNetNode,
"ACN_ReferenceControlNetFinetune": ReferenceControlFinetune,
# LOOSEControl
#"ACN_ControlNetLoaderWithLoraAdvanced": ControlNetLoaderWithLoraAdvanced,
# Deprecated
@@ -266,12 +267,13 @@ NODE_DISPLAY_NAME_MAPPINGS = {
# SparseCtrl
"ACN_SparseCtrlRGBPreprocessor": "RGB SparseCtrl 🛂🅐🅒🅝",
"ACN_SparseCtrlLoaderAdvanced": "Load SparseCtrl Model 🛂🅐🅒🅝",
"ACN_SparseCtrlMergedLoaderAdvanced": "Load Merged SparseCtrl Model 🛂🅐🅒🅝",
"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 🛂🅐🅒🅝",
"ACN_ReferenceControlNetFinetune": "Reference ControlNet (Finetune) 🛂🅐🅒🅝",
# LOOSEControl
#"ACN_ControlNetLoaderWithLoraAdvanced": "Load Adv. ControlNet Model w/ LoRA 🛂🅐🅒🅝",
# Deprecated
+31 -3
View File
@@ -24,9 +24,37 @@ class ReferenceControlNetNode:
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)
def load_controlnet(self, reference_type: str, style_fidelity: float, ref_weight: float):
ref_opts = ReferenceOptions.create_combo(reference_type=reference_type, style_fidelity=style_fidelity, ref_weight=ref_weight)
controlnet = ReferenceAdvanced(ref_opts=ref_opts, timestep_keyframes=None)
return (controlnet,)
class ReferenceControlFinetune:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"attn_style_fidelity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"attn_ref_weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"attn_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"adain_style_fidelity": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"adain_ref_weight": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"adain_strength": ("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,
attn_style_fidelity: float, attn_ref_weight: float, attn_strength: float,
adain_style_fidelity: float, adain_ref_weight: float, adain_strength: float):
ref_opts = ReferenceOptions(reference_type=ReferenceType.ATTN_ADAIN,
attn_style_fidelity=attn_style_fidelity, attn_ref_weight=attn_ref_weight, attn_strength=attn_strength,
adain_style_fidelity=adain_style_fidelity, adain_ref_weight=adain_ref_weight, adain_strength=adain_strength)
controlnet = ReferenceAdvanced(ref_opts=ref_opts, timestep_keyframes=None)
return (controlnet,)
-21
View File
@@ -316,27 +316,6 @@ def broadcast_image_to_full(tensor, target_batch_size, batched_number, except_on
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):
'''