Initial working implementation of ContextRef

This commit is contained in:
Jedrzej Kosinski
2024-07-27 21:08:46 -05:00
parent 4dde3062a0
commit ee3490d4b0
4 changed files with 65 additions and 24 deletions
+8 -5
View File
@@ -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):
+48 -16
View File
@@ -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
+5 -3
View File
@@ -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
+4
View File
@@ -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