Added keyframe support for ContextRef, small changes to timestep keyframe code, initial ContextRef version check
This commit is contained in:
@@ -14,7 +14,7 @@ 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,
|
||||
from .utils import (AdvancedControlBase, ControlWeights, TimestepKeyframeGroup, TimestepKeyframe, AbstractPreprocWrapper,
|
||||
broadcast_image_to_extend, ORIG_PREVIOUS_CONTROLNET, CONTROL_INIT_BY_ACN)
|
||||
|
||||
|
||||
@@ -38,6 +38,8 @@ CONTEXTREF_CONTROL_LIST_ALL = "contextref_control_list_all"
|
||||
CONTEXTREF_MACHINE_STATE = "contextref_machine_state"
|
||||
CONTEXTREF_TEMP_COND_IDX = "contextref_temp_cond_idx"
|
||||
|
||||
HIGHEST_VERSION_SUPPORT = 1
|
||||
|
||||
|
||||
class MachineState:
|
||||
WRITE = "write"
|
||||
@@ -281,9 +283,20 @@ class ReferenceAdvanced(ControlBase, AdvancedControlBase):
|
||||
|
||||
def handle_context_ref_setup(contextref_obj, transformer_options: dict, positive, negative):
|
||||
transformer_options[CONTEXTREF_MACHINE_STATE] = MachineState.OFF
|
||||
cref_opt_dict = contextref_obj.params.create_dict() # ContextRefParams obj from ADE
|
||||
# verify version is compatible
|
||||
if contextref_obj.version > HIGHEST_VERSION_SUPPORT:
|
||||
raise Exception(f"AnimateDiff-Evolved's ContextRef v{contextref_obj.version} is not supported in currently-installed Advanced-ControlNet (only supports ContextRef up to v{HIGHEST_VERSION_SUPPORT}); " +
|
||||
f"update your Advanced-ControlNet nodes for ContextRef to work.")
|
||||
# init ReferenceOptions
|
||||
cref_opt_dict = contextref_obj.tune.create_dict() # ContextRefTune obj from ADE
|
||||
opts = ReferenceOptions.create_from_kwargs(**cref_opt_dict)
|
||||
cref = ReferenceAdvanced(ref_opts=opts, timestep_keyframes=None)
|
||||
# init TimestepKeyframes
|
||||
cref_tks_list = contextref_obj.keyframe.create_list_of_dicts() # ContextRefKeyframeGroup obj from ADE
|
||||
timestep_keyframes = _create_tks_from_dict_list(cref_tks_list)
|
||||
# create ReferenceAdvanced
|
||||
cref = ReferenceAdvanced(ref_opts=opts, timestep_keyframes=timestep_keyframes)
|
||||
cref.strength = contextref_obj.strength # ContextRef obj from ADE
|
||||
cref.set_cond_hint_mask(contextref_obj.mask)
|
||||
cref.order = 99
|
||||
cref.is_context_ref = True
|
||||
context_ref_list = [cref]
|
||||
@@ -293,6 +306,27 @@ def handle_context_ref_setup(contextref_obj, transformer_options: dict, positive
|
||||
return context_ref_list
|
||||
|
||||
|
||||
def _create_tks_from_dict_list(dlist: list[dict[str]]) -> TimestepKeyframeGroup:
|
||||
tks = TimestepKeyframeGroup()
|
||||
if dlist is None or len(dlist) == 0:
|
||||
return tks
|
||||
for d in dlist:
|
||||
# scheduling
|
||||
start_percent = d["start_percent"]
|
||||
guarantee_steps = d["guarantee_steps"]
|
||||
inherit_missing = d["inherit_missing"]
|
||||
# values
|
||||
strength = d["strength"]
|
||||
mask = d["mask"]
|
||||
tune = d["tune"]
|
||||
mode = d["mode"]
|
||||
# create keyframe
|
||||
tk = TimestepKeyframe(start_percent=start_percent, guarantee_steps=guarantee_steps, inherit_missing=inherit_missing,
|
||||
strength=strength, mask_hint_orig=mask)
|
||||
tks.add(tk)
|
||||
return tks
|
||||
|
||||
|
||||
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
|
||||
|
||||
@@ -632,10 +632,8 @@ class AdvancedControlBase:
|
||||
self.weights = None
|
||||
self.latent_keyframes = None
|
||||
|
||||
def prepare_current_timestep(self, t: Tensor, batched_number: int):
|
||||
def prepare_current_timestep(self, t: Tensor, batched_number: int=1):
|
||||
self.t = float(t[0])
|
||||
self.batched_number = batched_number
|
||||
self.batch_size = len(t)
|
||||
# check if t has changed (otherwise do nothing, as step already accounted for)
|
||||
if self.t == self.prev_t:
|
||||
return
|
||||
@@ -749,6 +747,8 @@ class AdvancedControlBase:
|
||||
return True
|
||||
|
||||
def get_control_inject(self, x_noisy, t, cond, batched_number):
|
||||
self.batched_number = batched_number
|
||||
self.batch_size = len(t)
|
||||
# prepare timestep and everything related
|
||||
self.prepare_current_timestep(t=t, batched_number=batched_number)
|
||||
# if should not perform any actions for the controlnet, exit without doing any work
|
||||
|
||||
Reference in New Issue
Block a user