Added sliding reference window mode for ContextRef + ContextRef Mode nodes
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user