From cb309c60784700155afd0eb8cef6acca75d38094 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 1 Apr 2024 03:44:11 -0500 Subject: [PATCH 1/5] Added initial reference_adain and reference_adain+attn support --- adv_control/control_reference.py | 235 +++++++++++++++++++++++++++++-- adv_control/utils.py | 21 --- 2 files changed, 220 insertions(+), 36 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index c24932f..4838c3b 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -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 @@ -106,7 +145,8 @@ comfy.sample.sample_custom = refcn_sample_factory(comfy.sample.sample_custom, is 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_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 +155,7 @@ class MachineState: WRITE = "write" READ = "read" STYLEALIGN = "stylealign" - TEST = "test" + OFF = "off" class ReferenceType: @@ -124,8 +164,17 @@ 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: @@ -170,7 +219,7 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): 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() if is_mid: @@ -184,6 +233,14 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): 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() + b, c, h, w = x.shape + real_mask = torch.ones([b, 1, h, w]).to(dtype=x.dtype, device=x.device) * self.strength + self.apply_advanced_strengths_and_masks(x=real_mask, batched_number=self.batched_number) + 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: @@ -225,7 +282,7 @@ 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)) # prepare mask - use direct_attn, so the mask dims will match source latents (and be smaller) @@ -258,6 +315,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 +355,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 +397,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 +426,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,9 +464,26 @@ 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 if any ref_controlnets will use adain + use_adain = False + for control in ref_controlnets: + if ReferenceType.is_adain(control.ref_opts.reference_type): + use_adain = True + break + if use_adain: + # 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 + 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_CONTROL_LIST] = [control] # handle masks - apply x to unmasked #strength_mask = torch.ones_like(x, dtype=x.dtype) * control.strength @@ -348,12 +496,15 @@ 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_ATTN_MACHINE_STATE] = MachineState.READ + transformer_options[REF_ADAIN_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() + if use_adain: + openaimodel.forward_timestep_embed = orig_forward_timestep_embed return forward_inject_UNetModel @@ -398,7 +549,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten 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_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: @@ -435,7 +586,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten 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) + effective_strength = ref_controlnets[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, @@ -466,7 +617,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten 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) + effective_strength = ref_controlnets[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 +689,60 @@ 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_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.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.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: + # TODO: support strength/masks/latent_kfs + bank_styles = ts.injection_holder.bank_styles + style_fidelity = bank_styles.get_avg_style_fidelity() + var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) + std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 + var_acc = bank_styles.get_avg_var_bank() + mean_acc = bank_styles.get_avg_mean_bank() + std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 + y_uc = (((x - mean) / std) * std_acc) + mean_acc + if ref_controlnets[0].any_strength_to_apply(): + effective_strength = ref_controlnets[0].get_effective_adain_mask_or_float(x=x) + y_uc = y_uc * effective_strength + x * (1-effective_strength) + 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] diff --git a/adv_control/utils.py b/adv_control/utils.py index 5caefbd..e2fd229 100644 --- a/adv_control/utils.py +++ b/adv_control/utils.py @@ -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): ''' From 414ac8d6ec5b36812395a12a5ce543091bc32832 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 1 Apr 2024 04:28:00 -0500 Subject: [PATCH 2/5] Started separating out attn and adain variables to prevent issues when using multiple RefCN --- adv_control/control_reference.py | 29 +++++++++++++++++------------ 1 file changed, 17 insertions(+), 12 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 4838c3b..9e4998f 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -142,7 +142,8 @@ 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_ATTN_MACHINE_STATE = "ref_attn_machine_state" @@ -464,13 +465,15 @@ 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 if any ref_controlnets will use adain - use_adain = False + # 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): - use_adain = True - break - if use_adain: + 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) @@ -484,7 +487,8 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): transformer_options[REF_ADAIN_MACHINE_STATE] = MachineState.WRITE else: transformer_options[REF_ADAIN_MACHINE_STATE] = MachineState.OFF - transformer_options[REF_CONTROL_LIST] = [control] + transformer_options[REF_ATTN_CONTROL_LIST] = [control] + transformer_options[REF_ADAIN_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) @@ -498,12 +502,13 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): # run diffusion for real now transformer_options[REF_ATTN_MACHINE_STATE] = MachineState.READ transformer_options[REF_ADAIN_MACHINE_STATE] = MachineState.READ - transformer_options[REF_CONTROL_LIST] = ref_controlnets + 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 use_adain: + if len(adain_controlnets) > 0: openaimodel.forward_timestep_embed = orig_forward_timestep_embed return forward_inject_UNetModel @@ -548,7 +553,7 @@ 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_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: @@ -705,7 +710,7 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory 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_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 @@ -721,9 +726,9 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory if len(ts.injection_holder.bank_styles.var_bank) > 0: # TODO: support strength/masks/latent_kfs bank_styles = ts.injection_holder.bank_styles - style_fidelity = bank_styles.get_avg_style_fidelity() var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 + style_fidelity = bank_styles.get_avg_style_fidelity() var_acc = bank_styles.get_avg_var_bank() mean_acc = bank_styles.get_avg_mean_bank() std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 From 5e3c34027f60da1bd205651c43587a80e3854726 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 1 Apr 2024 18:20:35 -0500 Subject: [PATCH 3/5] Added Reference ControlNet (Finetune) to control attn and adain values separately --- adv_control/control_reference.py | 79 ++++++++++++++++++++++---------- adv_control/nodes.py | 6 ++- adv_control/nodes_reference.py | 35 ++++++++++++-- 3 files changed, 91 insertions(+), 29 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 9e4998f..2f5e7a7 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -179,17 +179,40 @@ class ReferenceType: 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." @@ -207,27 +230,31 @@ 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_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 @@ -236,9 +263,9 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): def get_effective_adain_mask_or_float(self, x: Tensor): if not self.should_apply_effective_masks: - return self.get_effective_strength() + 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 + 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 @@ -250,9 +277,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 @@ -285,7 +314,8 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): self.cond_hint = self.latent_format.process_in(self.cond_hint) 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 @@ -300,7 +330,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): @@ -557,9 +588,9 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten 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: @@ -590,7 +621,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten 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(): + if ref_controlnets[idx].any_attn_strength_to_apply(): effective_strength = ref_controlnets[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]( @@ -621,7 +652,7 @@ def _forward_inject_BasicTransformerBlock(self: RefBasicTransformerBlock, x: Ten 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(): + if ref_controlnets[idx].any_attn_strength_to_apply(): effective_strength = ref_controlnets[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( @@ -715,11 +746,11 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory # if in WRITE mode, save var, mean, and style_cfg if ref_machine_state == MachineState.WRITE: - if ref_controlnets[0].ref_opts.ref_weight > ts.injection_holder.gn_weight: + 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.style_fidelity) + 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: @@ -733,7 +764,7 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory mean_acc = bank_styles.get_avg_mean_bank() std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 y_uc = (((x - mean) / std) * std_acc) + mean_acc - if ref_controlnets[0].any_strength_to_apply(): + if ref_controlnets[0].any_adain_strength_to_apply(): effective_strength = ref_controlnets[0].get_effective_adain_mask_or_float(x=x) y_uc = y_uc * effective_strength + x * (1-effective_strength) y_c = y_uc.clone() diff --git a/adv_control/nodes.py b/adv_control/nodes.py index cea6ada..3544a30 100644 --- a/adv_control/nodes.py +++ b/adv_control/nodes.py @@ -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 diff --git a/adv_control/nodes_reference.py b/adv_control/nodes_reference.py index e51694b..5556cf3 100644 --- a/adv_control/nodes_reference.py +++ b/adv_control/nodes_reference.py @@ -24,9 +24,38 @@ 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": { + "reference_type": (ReferenceType._LIST, {"default": ReferenceType.ATTN_ADAIN}), + "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, reference_type: str, + 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=reference_type, + 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,) From 019452c790968a2b9c2bd99c34422fcf355a1fa1 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 2 Apr 2024 01:08:19 -0500 Subject: [PATCH 4/5] Made sure ref cns match up to expected bank indexes, made adain properly react to strengths and masks --- adv_control/control_reference.py | 77 +++++++++++++++++++++++--------- 1 file changed, 57 insertions(+), 20 deletions(-) diff --git a/adv_control/control_reference.py b/adv_control/control_reference.py index 2f5e7a7..8a585a1 100644 --- a/adv_control/control_reference.py +++ b/adv_control/control_reference.py @@ -269,6 +269,20 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase): 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: @@ -353,9 +367,9 @@ def ref_noise_latents(latents: Tensor, sigma: Tensor, noise: Tensor=None): 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(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) @@ -520,10 +534,7 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections): transformer_options[REF_ADAIN_MACHINE_STATE] = MachineState.OFF transformer_options[REF_ATTN_CONTROL_LIST] = [control] transformer_options[REF_ADAIN_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() @@ -620,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_attn_strength_to_apply(): - effective_strength = ref_controlnets[idx].get_effective_attn_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, @@ -651,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_attn_strength_to_apply(): - effective_strength = ref_controlnets[idx].get_effective_attn_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, @@ -755,18 +780,30 @@ def forward_timestep_embed_ref_inject_factory(orig_timestep_embed_inject_factory # 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: - # TODO: support strength/masks/latent_kfs 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 - style_fidelity = bank_styles.get_avg_style_fidelity() - var_acc = bank_styles.get_avg_var_bank() - mean_acc = bank_styles.get_avg_mean_bank() - std_acc = torch.maximum(var_acc, torch.zeros_like(var_acc) + eps) ** 0.5 - y_uc = (((x - mean) / std) * std_acc) + mean_acc - if ref_controlnets[0].any_adain_strength_to_apply(): - effective_strength = ref_controlnets[0].get_effective_adain_mask_or_float(x=x) - y_uc = y_uc * effective_strength + x * (1-effective_strength) + 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] From 0357841c15810d3527907570542f4e1434017e49 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 2 Apr 2024 01:14:06 -0500 Subject: [PATCH 5/5] Removed unnecessary dropdown from Reference ControlNet (Finetune) node, as should always be adain+attn anyway to prevent confusion --- adv_control/nodes_reference.py | 5 ++--- 1 file changed, 2 insertions(+), 3 deletions(-) diff --git a/adv_control/nodes_reference.py b/adv_control/nodes_reference.py index 5556cf3..fd0e4dc 100644 --- a/adv_control/nodes_reference.py +++ b/adv_control/nodes_reference.py @@ -35,7 +35,6 @@ class ReferenceControlFinetune: def INPUT_TYPES(s): return { "required": { - "reference_type": (ReferenceType._LIST, {"default": ReferenceType.ATTN_ADAIN}), "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}), @@ -50,10 +49,10 @@ class ReferenceControlFinetune: CATEGORY = "Adv-ControlNet ๐Ÿ›‚๐Ÿ…๐Ÿ…’๐Ÿ…/Reference" - def load_controlnet(self, reference_type: str, + 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=reference_type, + 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)