Added sliding reference window mode for ContextRef + ContextRef Mode nodes

This commit is contained in:
Jedrzej Kosinski
2024-07-31 06:31:21 -05:00
parent 1c68c56809
commit 11252fc34f
5 changed files with 104 additions and 20 deletions
+20 -1
View File
@@ -46,10 +46,29 @@ class ContextRefParams:
self.adain_strength = adain_strength
class ContextRefMode:
FIRST = "first"
SLIDING = "sliding"
_LIST = [FIRST, SLIDING]
def __init__(self, mode: str, sliding_width=2):
self.mode = mode
self.sliding_width = sliding_width
@classmethod
def init_first(cls):
return ContextRefMode(cls.FIRST)
@classmethod
def init_sliding(cls, sliding_width):
return ContextRefMode(cls.SLIDING, sliding_width=sliding_width)
class ContextRef(ContextExtra):
def __init__(self, start_percent: float, end_percent: float, params: ContextRefParams):
def __init__(self, start_percent: float, end_percent: float, params: ContextRefParams, mode: ContextRefMode):
super().__init__(start_percent=start_percent, end_percent=end_percent)
self.params = params
self.mode = mode
def should_run(self):
return super().should_run()
+6 -1
View File
@@ -29,7 +29,8 @@ from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, Weig
from .nodes_context import (LegacyLoopedUniformContextOptionsNode, LoopedUniformContextOptionsNode, LoopedUniformViewOptionsNode, StandardUniformContextOptionsNode, StandardStaticContextOptionsNode, BatchedContextOptionsNode,
StandardStaticViewOptionsNode, StandardUniformViewOptionsNode, ViewAsContextOptionsNode,
VisualizeContextOptionsK, VisualizeContextOptionsKAdv, VisualizeContextOptionsSCustom,
SetContextExtrasOnContextOptions, ContextExtras_NaiveReuse, ContextExtras_ContextRef)
SetContextExtrasOnContextOptions, ContextExtras_NaiveReuse, ContextExtras_ContextRef,
ContextRef_ModeFirst, ContextRef_ModeSliding)
from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, SweetspotStretchPENode, FullStretchPENode,
WeightAdjustAllAddNode, WeightAdjustAllMultNode, WeightAdjustIndivAddNode, WeightAdjustIndivMultNode,
WeightAdjustIndivAttnAddNode, WeightAdjustIndivAttnMultNode)
@@ -75,6 +76,8 @@ NODE_CLASS_MAPPINGS = {
"ADE_ContextExtras_Set": SetContextExtrasOnContextOptions,
"ADE_ContextExtras_ContextRef": ContextExtras_ContextRef,
"ADE_ContextExtras_NaiveReuse": ContextExtras_NaiveReuse,
"ADE_ContextExtras_ContextRef_ModeFirst": ContextRef_ModeFirst,
"ADE_ContextExtras_ContextRef_ModeSliding": ContextRef_ModeSliding,
#------------------------------------------------------------------------------
###############################################################################
# Iteration Opts
@@ -213,6 +216,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"ADE_ContextExtras_Set": "Set Context Extras 🎭🅐🅓",
"ADE_ContextExtras_ContextRef": "Context Extras◆ContextRef 🎭🅐🅓",
"ADE_ContextExtras_NaiveReuse": "Context Extras◆NaiveReuse 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_ModeFirst": "ContextRef Mode◆First 🎭🅐🅓",
"ADE_ContextExtras_ContextRef_ModeSliding": "ContextRef Mode◆Sliding 🎭🅐🅓",
#------------------------------------------------------------------------------
###############################################################################
# Iteration Opts
+47 -2
View File
@@ -7,7 +7,7 @@ from comfy.model_patcher import ModelPatcher
from .context import (ContextFuseMethod, ContextOptions, ContextOptionsGroup, ContextSchedules,
generate_context_visualization)
from .context_extras import ContextExtrasGroup, ContextRef, ContextRefParams, NaiveReuse
from .context_extras import ContextExtrasGroup, ContextRef, ContextRefParams, ContextRefMode, NaiveReuse
from .utils_model import BIGMAX, MAX_RESOLUTION
@@ -510,6 +510,7 @@ class ContextExtras_ContextRef:
"optional": {
"prev_extras": ("CONTEXT_EXTRAS",),
"strength_multival": ("MULTIVAL",),
"contextref_mode": ("CONTEXTREF_MODE",),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
"end_percent": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.001}),
"autosize": ("ADEAUTOSIZE", {"padding": 0}),
@@ -521,6 +522,7 @@ class ContextExtras_ContextRef:
FUNCTION = "create_context_extra"
def create_context_extra(self, start_percent=0.0, end_percent=0.1, strength_multival: Union[float, Tensor]=None,
contextref_mode: ContextRefMode=None,
prev_extras: ContextExtrasGroup=None):
if prev_extras is None:
prev_extras = prev_extras = ContextExtrasGroup()
@@ -528,6 +530,49 @@ class ContextExtras_ContextRef:
# create extra
# TODO: make customizable, and allow mask input
params = ContextRefParams(attn_style_fidelity=1.0, attn_ref_weight=1.0, attn_atrength=1.0)
context_ref = ContextRef(start_percent=start_percent, end_percent=end_percent, params=params)
if contextref_mode is None:
contextref_mode = ContextRefMode.init_first()
context_ref = ContextRef(start_percent=start_percent, end_percent=end_percent, params=params, mode=contextref_mode)
prev_extras.add(context_ref)
return (prev_extras,)
class ContextRef_ModeFirst:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
},
"optional": {
"autosize": ("ADEAUTOSIZE", {"padding": 25}),
},
}
RETURN_TYPES = ("CONTEXTREF_MODE",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
FUNCTION = "create_contextref_mode"
def create_contextref_mode(self):
mode = ContextRefMode.init_first()
return (mode,)
class ContextRef_ModeSliding:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
},
"optional": {
"sliding_width": ("INT", {"default": 2, "min": 2, "max": BIGMAX, "step": 1}),
"autosize": ("ADEAUTOSIZE", {"padding": 42}),
}
}
RETURN_TYPES = ("CONTEXTREF_MODE",)
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras/ContextRef"
FUNCTION = "create_contextref_mode"
def create_contextref_mode(self, sliding_width):
mode = ContextRefMode.init_sliding(sliding_width=sliding_width)
return (mode,)
+25 -16
View File
@@ -26,8 +26,9 @@ import comfy.ops
from .conditioning import COND_CONST, LoraHookGroup, conditioning_set_values
from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows
from .context_extras import ContextRefMode
from .sample_settings import IterationOptions, SampleSettings, SeedNoiseGeneration, NoisedImageToInject
from .utils_model import ModelTypeSD, vae_encode_raw_batched, vae_decode_raw_batched
from .utils_model import ModelTypeSD, MachineState, vae_encode_raw_batched, vae_decode_raw_batched
from .utils_motion import composite_extend, get_combined_multival, prepare_mask_batch, extend_to_batch_size
from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher
from .motion_module_ad import AnimateDiffFormat, AnimateDiffInfo, AnimateDiffVersion, VanillaTemporalModule
@@ -831,27 +832,31 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options
CONTEXTREF_CONTROL_LIST_ALL = "contextref_control_list_all"
CONTEXTREF_MACHINE_STATE = "contextref_machine_state"
CONTEXTREF_CLEAN_FUNC = "contextref_clean_func"
context_ref = False
context_ref_injector = None
contextref_active = False
contextref_injector = None
contextref_mode = None
first_context = False
# need to make sure that contextref stuff gets cleaned up, no matter what
try:
if ADGS.params.context_options.extras.should_run_context_ref():
context_ref = True
contextref_active = True
first_context = True
contextref_mode = ADGS.params.context_options.extras.context_ref.mode
# use injector to ensure only 1 cond or uncond will be batched at a time
context_ref_injector = ContextRefInjector()
context_ref_injector.inject()
contextref_injector = ContextRefInjector()
contextref_injector.inject()
naive_init = False
curr_window_idx = -1
naivereuse_active = False
cached_naive_conds = None
cached_naive_ctx_idxs = None
if ADGS.params.context_options.extras.should_run_naive_reuse():
cached_naive_conds = [torch.zeros_like(x_in) for _ in conds]
#cached_naive_counts = [torch.zeros((x_in.shape[0], 1, 1, 1), device=x_in.device) for _ in conds]
naive_init = True
naivereuse_active = True
# perform calc_conds_batch per context window
for ctx_idxs in context_windows:
curr_window_idx += 1
ADGS.params.sub_idxs = ctx_idxs
if ADGS.motion_models is not None:
ADGS.motion_models.set_sub_idxs(ctx_idxs)
@@ -869,16 +874,20 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options
sub_timestep = timestep[full_idxs]
sub_conds = [get_resized_cond(cond, full_idxs, len(ctx_idxs)) for cond in conds]
if context_ref:
if contextref_active:
# set cond counter to 0 (each cond encountered will increment it by 1)
model_options["transformer_options"][CONTEXTREF_CONTROL_LIST_ALL][0].contextref_cond_idx = 0
if first_context:
first_context = False
model_options["transformer_options"][CONTEXTREF_MACHINE_STATE] = "write"
model_options["transformer_options"][CONTEXTREF_MACHINE_STATE] = MachineState.WRITE
else:
model_options["transformer_options"][CONTEXTREF_MACHINE_STATE] = "read"
model_options["transformer_options"][CONTEXTREF_MACHINE_STATE] = MachineState.READ
if contextref_mode.mode == ContextRefMode.SLIDING: # if sliding, check if time to READ and WRITE
if curr_window_idx % (contextref_mode.sliding_width-1) == 0:
model_options["transformer_options"][CONTEXTREF_MACHINE_STATE] = MachineState.READ_WRITE
else:
model_options["transformer_options"][CONTEXTREF_MACHINE_STATE] = "off"
model_options["transformer_options"][CONTEXTREF_MACHINE_STATE] = MachineState.OFF
#logger.info(f"window: {curr_window_idx} - {model_options['transformer_options'][CONTEXTREF_MACHINE_STATE]}")
sub_conds_out = calc_cond_uncond_batch_wrapper(model, sub_conds, sub_x, sub_timestep, model_options)
@@ -904,16 +913,16 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options
for i in range(len(sub_conds_out)):
conds_final[i][full_idxs] += sub_conds_out[i] * weights_tensor
counts_final[i][full_idxs] += weights_tensor
if naive_init:
if naivereuse_active:
cached_naive_ctx_idxs = ctx_idxs
for i in range(len(sub_conds)):
cached_naive_conds[i][full_idxs] = conds_final[i][full_idxs] / counts_final[i][full_idxs]
naive_init = False
naivereuse_active = False
finally:
# clean contextref stuff with provided ACN function, if applicable
if context_ref:
if contextref_active:
model_options["transformer_options"][CONTEXTREF_CLEAN_FUNC]()
context_ref_injector.restore()
contextref_injector.restore()
# handle NaiveReuse
if cached_naive_conds is not None:
+6
View File
@@ -26,6 +26,12 @@ BIGMAX = (2**53-1)
MAX_RESOLUTION = 16384 # mirrors ComfyUI's nodes.py MAX_RESOLUTION
class MachineState:
READ = "read"
WRITE = "write"
READ_WRITE = "read_write"
OFF = "off"
def vae_encode_raw_dynamic_batched(vae: VAE, pixels: Tensor, max_batch=16, min_batch=1, max_size=512*512, show_pbar=False):
b, h, w, c = pixels.shape