diff --git a/animatediff/context_extras.py b/animatediff/context_extras.py index 1d28edb..350ce9c 100644 --- a/animatediff/context_extras.py +++ b/animatediff/context_extras.py @@ -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() diff --git a/animatediff/nodes.py b/animatediff/nodes.py index e14e9ca..4796f90 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -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 diff --git a/animatediff/nodes_context.py b/animatediff/nodes_context.py index 8f69aad..37b6c46 100644 --- a/animatediff/nodes_context.py +++ b/animatediff/nodes_context.py @@ -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,) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index e8b294b..87aedf4 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -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: diff --git a/animatediff/utils_model.py b/animatediff/utils_model.py index 0e62d00..48c432b 100644 --- a/animatediff/utils_model.py +++ b/animatediff/utils_model.py @@ -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