Initial working implementation of ContextRef
This commit is contained in:
@@ -16,13 +16,10 @@ from .control_lllite import LLLiteModule, LLLitePatch, load_controllllite
|
||||
from .control_svd import svd_unet_config_from_diffusers_unet, SVDControlNet, svd_unet_to_diffusers
|
||||
from .utils import (AdvancedControlBase, TimestepKeyframeGroup, LatentKeyframeGroup, AbstractPreprocWrapper, ControlWeightType, ControlWeights, WeightTypeException,
|
||||
manual_cast_clean_groupnorm, disable_weight_init_clean_groupnorm, prepare_mask_batch, get_properly_arranged_t2i_weights, load_torch_file_with_dict_factory,
|
||||
broadcast_image_to_extend, extend_to_batch_size)
|
||||
broadcast_image_to_extend, extend_to_batch_size, ORIG_PREVIOUS_CONTROLNET, CONTROL_INIT_BY_ACN)
|
||||
from .logger import logger
|
||||
|
||||
|
||||
ORIG_PREVIOUS_CONTROLNET = "_orig_previous_controlnet"
|
||||
|
||||
|
||||
class ControlNetAdvanced(ControlNet, AdvancedControlBase):
|
||||
def __init__(self, control_model, timestep_keyframes: TimestepKeyframeGroup, global_average_pooling=False, compression_ratio=8, latent_format=None, device=None, load_device=None, manual_cast_dtype=None):
|
||||
super().__init__(control_model=control_model, global_average_pooling=global_average_pooling, compression_ratio=compression_ratio, latent_format=latent_format, device=device, load_device=load_device, manual_cast_dtype=manual_cast_dtype)
|
||||
@@ -574,7 +571,13 @@ def restore_all_controlnet_conns(conds: list[list[dict[str]]]):
|
||||
if main_cond is not None:
|
||||
for cond in main_cond:
|
||||
if "control" in cond[1]:
|
||||
_restore_all_controlnet_conns(cond[1]["control"])
|
||||
# if ACN is the one to have initialized it, delete it
|
||||
# TODO: maybe check if someone else did a similar hack, and carefully pluck out our stuff?
|
||||
if CONTROL_INIT_BY_ACN in cond[1]:
|
||||
cond[1].pop("control")
|
||||
cond[1].pop(CONTROL_INIT_BY_ACN)
|
||||
else:
|
||||
_restore_all_controlnet_conns(cond[1]["control"])
|
||||
|
||||
|
||||
def _restore_all_controlnet_conns(input_object: ControlBase):
|
||||
|
||||
@@ -14,7 +14,7 @@ from comfy.ldm.modules.diffusionmodules import openaimodel
|
||||
|
||||
from .logger import logger
|
||||
from .utils import (AdvancedControlBase, ControlWeights, TimestepKeyframeGroup, AbstractPreprocWrapper,
|
||||
broadcast_image_to_extend)
|
||||
broadcast_image_to_extend, ORIG_PREVIOUS_CONTROLNET, CONTROL_INIT_BY_ACN)
|
||||
|
||||
|
||||
REF_READ_ATTN_CONTROL_LIST = "ref_read_attn_control_list"
|
||||
@@ -255,18 +255,50 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
|
||||
return self
|
||||
|
||||
|
||||
def handle_context_ref_setup(transformer_options):
|
||||
transformer_options[CONTEXTREF_ATTN_MACHINE_STATE] = MachineState.OFF
|
||||
transformer_options[CONTEXTREF_ADAIN_MACHINE_STATE] = MachineState.OFF
|
||||
def handle_context_ref_setup(transformer_options, positive, negative):
|
||||
transformer_options[CONTEXTREF_MACHINE_STATE] = MachineState.OFF
|
||||
opts = ReferenceOptions(ReferenceType.ATTN, attn_style_fidelity=1.0, attn_ref_weight=1.0, attn_strength=1.0, adain_style_fidelity=0.0, adain_ref_weight=0.0)
|
||||
cref = ReferenceAdvanced(ref_opts=opts, timestep_keyframes=None)
|
||||
cref.order = -1
|
||||
cref.order = 99
|
||||
cref.is_context_ref = True
|
||||
context_ref_list = [cref]
|
||||
transformer_options[CONTEXTREF_CONTROL_LIST_ALL] = context_ref_list
|
||||
transformer_options[CONTEXTREF_OPTIONS_CLASS] = ReferenceOptions
|
||||
_add_context_ref_to_conds([positive, negative], cref)
|
||||
return context_ref_list
|
||||
|
||||
|
||||
def _add_context_ref_to_conds(conds: list[list[dict[str]]], context_ref: ReferenceAdvanced):
|
||||
def _add_context_ref_to_existing_control(control: ControlBase, context_ref: ReferenceAdvanced):
|
||||
curr_cn = control
|
||||
while curr_cn is not None:
|
||||
if type(curr_cn) == ReferenceAdvanced and curr_cn.is_context_ref:
|
||||
break
|
||||
if curr_cn.previous_controlnet is not None:
|
||||
curr_cn = curr_cn.previous_controlnet
|
||||
continue
|
||||
orig_previous_controlnet = curr_cn.previous_controlnet
|
||||
# NOTE: code is already in place to restore any ORIG_PREVIOUS_CONTROLNET props
|
||||
setattr(curr_cn, ORIG_PREVIOUS_CONTROLNET, orig_previous_controlnet)
|
||||
curr_cn.previous_controlnet = context_ref
|
||||
curr_cn = orig_previous_controlnet
|
||||
|
||||
def _add_context_ref(actual_cond: dict[str], context_ref: ReferenceAdvanced):
|
||||
# if controls already present on cond, add it to the last previous_controlnet
|
||||
if "control" in actual_cond:
|
||||
return _add_context_ref_to_existing_control(actual_cond["control"], context_ref)
|
||||
# otherwise, need to add it to begin with, and should mark that it should be cleaned after
|
||||
actual_cond["control"] = context_ref
|
||||
actual_cond[CONTROL_INIT_BY_ACN] = True
|
||||
|
||||
# either add context_ref to end of existing cnet chain, or init 'control' key on actual cond
|
||||
for cond in conds:
|
||||
if cond is not None:
|
||||
for sub_cond in cond:
|
||||
actual_cond = sub_cond[1]
|
||||
_add_context_ref(actual_cond, context_ref)
|
||||
|
||||
|
||||
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)
|
||||
@@ -714,29 +746,29 @@ def factory_forward_inject_UNetModel(reference_injections: ReferenceInjections):
|
||||
|
||||
# do contextref stuff, if needed
|
||||
if len(context_controlnets) > 0:
|
||||
# TODO: clean contextref stuff if attn writing or off
|
||||
if transformer_options[CONTEXTREF_ATTN_MACHINE_STATE] in [MachineState.WRITE, MachineState.OFF]:
|
||||
# TODO: clean contextref stuff if contextref is writing or off
|
||||
if transformer_options[CONTEXTREF_MACHINE_STATE] in [MachineState.WRITE, MachineState.OFF]:
|
||||
reference_injections.clean_contextref_module_mem()
|
||||
### add ContextRef to appropriate lists
|
||||
# attn
|
||||
if transformer_options[CONTEXTREF_ATTN_MACHINE_STATE] == MachineState.WRITE:
|
||||
write_attn_list.extend(context_attn_controlnets)
|
||||
elif transformer_options[CONTEXTREF_ATTN_MACHINE_STATE] == MachineState.READ:
|
||||
if transformer_options[CONTEXTREF_MACHINE_STATE] == MachineState.READ:
|
||||
read_attn_list.extend(context_attn_controlnets)
|
||||
elif transformer_options[CONTEXTREF_MACHINE_STATE] == MachineState.WRITE:
|
||||
write_attn_list.extend(context_attn_controlnets)
|
||||
# adain
|
||||
if transformer_options[CONTEXTREF_ADAIN_MACHINE_STATE] == MachineState.WRITE:
|
||||
write_attn_list.extend(context_adain_controlnets)
|
||||
elif transformer_options[CONTEXTREF_ADAIN_MACHINE_STATE] == MachineState.READ:
|
||||
read_attn_list.extend(context_adain_controlnets)
|
||||
if transformer_options[CONTEXTREF_MACHINE_STATE] == MachineState.READ:
|
||||
read_adain_list.extend(context_adain_controlnets)
|
||||
elif transformer_options[CONTEXTREF_MACHINE_STATE] == MachineState.WRITE:
|
||||
write_adain_list.extend(context_adain_controlnets)
|
||||
# apply lists, containing both RefCN and ContextRef
|
||||
transformer_options[REF_READ_ATTN_CONTROL_LIST] = read_attn_list
|
||||
transformer_options[REF_WRITE_ATTN_CONTROL_LIST] = write_attn_list
|
||||
transformer_options[REF_READ_ADAIN_CONTROL_LIST] = read_adain_list
|
||||
transformer_options[REF_WRITE_ADAIN_CONTROL_LIST] = write_adain_list
|
||||
|
||||
# run diffusion for real
|
||||
return reference_injections.diffusion_model_orig_forward(x, *args, **kwargs)
|
||||
finally:
|
||||
# make sure banks are cleared no matter what happens - otherwise, RIP VRAM
|
||||
# make sure ref banks are cleared no matter what happens - otherwise, RIP VRAM
|
||||
reference_injections.clean_ref_module_mem()
|
||||
if len(adain_controlnets) > 0:
|
||||
openaimodel.forward_timestep_embed = orig_forward_timestep_embed
|
||||
|
||||
@@ -49,7 +49,7 @@ def acn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable
|
||||
ref_set: set[ReferenceAdvanced] = set()
|
||||
if control is None:
|
||||
return ref_set
|
||||
if type(control) == ReferenceAdvanced:
|
||||
if type(control) == ReferenceAdvanced and not control.is_context_ref:
|
||||
control.order = order
|
||||
order -= 1
|
||||
ref_set.add(control)
|
||||
@@ -79,8 +79,6 @@ def acn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable
|
||||
if has_sliding_context_windows(model):
|
||||
model.model_options = model.model_options.copy()
|
||||
model.model_options["transformer_options"] = model.model_options["transformer_options"].copy()
|
||||
if has_contextref_enabled(model):
|
||||
context_refs = handle_context_ref_setup(model.model_options["transformer_options"])
|
||||
# convert all CNs to Advanced if needed
|
||||
controlnets_modified, positive, negative = support_sliding_context_windows(model, positive, negative)
|
||||
if controlnets_modified:
|
||||
@@ -88,6 +86,10 @@ def acn_sample_factory(orig_comfy_sample: Callable, is_custom=False) -> Callable
|
||||
args[-3] = positive
|
||||
args[-2] = negative
|
||||
args = tuple(args)
|
||||
# enable ContextRef, if requested
|
||||
if has_contextref_enabled(model):
|
||||
context_refs = handle_context_ref_setup(model.model_options["transformer_options"], positive, negative)
|
||||
controlnets_modified = True
|
||||
# look for Advanced ControlNets that will require intervention to work
|
||||
ref_set = set()
|
||||
lllite_dict: dict[ControlLLLiteAdvanced, None] = {} # dicts preserve insertion order since py3.7
|
||||
|
||||
@@ -21,6 +21,10 @@ from .logger import logger
|
||||
BIGMIN = -(2**53-1)
|
||||
BIGMAX = (2**53-1)
|
||||
|
||||
ORIG_PREVIOUS_CONTROLNET = "_orig_previous_controlnet"
|
||||
CONTROL_INIT_BY_ACN = "_control_init_by_ACN"
|
||||
|
||||
|
||||
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):
|
||||
# immediately restore load_torch_file to original version
|
||||
|
||||
Reference in New Issue
Block a user