diff --git a/animatediff/context.py b/animatediff/context.py index 50af0d8..c0514f0 100644 --- a/animatediff/context.py +++ b/animatediff/context.py @@ -90,6 +90,7 @@ class ContextOptionsGroup: self._current_index = 0 self.step = 0 self._set_first_as_current() + self.extras.cleanup() @property def step(self): diff --git a/animatediff/context_extras.py b/animatediff/context_extras.py index 519e0e3..1d28edb 100644 --- a/animatediff/context_extras.py +++ b/animatediff/context_extras.py @@ -2,6 +2,8 @@ from torch import Tensor from comfy.model_base import BaseModel +from .utils_motion import prepare_mask_batch, extend_to_batch_size, get_combined_multival + class ContextExtra: def __init__(self, start_percent: float, end_percent: float): @@ -24,9 +26,12 @@ class ContextExtra: return False return True + def cleanup(self): + pass + ################################ -# Context Ref +# ContextRef class ContextRefParams: def __init__(self, attn_style_fidelity=0.0, attn_ref_weight=0.0, attn_atrength=0.0, @@ -54,11 +59,30 @@ class ContextRef(ContextExtra): ################################ # NaiveReuse class NaiveReuse(ContextExtra): - def __init__(self, start_percent: float, end_percent: float, weighted_mean: float, mask_opt: Tensor=None): + def __init__(self, start_percent: float, end_percent: float, weighted_mean: float, multival_opt: Tensor=None): super().__init__(start_percent=start_percent, end_percent=end_percent) self.weighted_mean = weighted_mean - self.mask_opt = mask_opt + self.orig_multival = multival_opt + self.mask: Tensor = None + def cleanup(self): + super().cleanup() + del self.mask + self.mask = None + + def get_effective_weighted_mean(self, x: Tensor, idxs: list[int]): + if self.orig_multival is None: + return self.weighted_mean + # otherwise, is Tensor and should be extended to match dims and size of x; + # see if needs to be recalculated + if type(self.orig_multival) != Tensor: + return self.weighted_mean * self.orig_multival + elif self.mask is None or self.mask.shape[0] != x.shape[0] or self.mask.shape[-1] != x.shape[-1] or self.mask.shape[-2] != x.shape[-2]: + del self.mask + self.mask = prepare_mask_batch(self.orig_multival, x.shape) + self.mask = extend_to_batch_size(self.mask, x.shape[0]) + return self.weighted_mean * self.mask[idxs].to(dtype=x.dtype, device=x.device) + def should_run(self): to_return = super().should_run() # if weighted_mean is 0.0, then reuse will take no effect anyway @@ -105,6 +129,10 @@ class ContextExtrasGroup: else: raise Exception(f"Unrecognized ContextExtras type: {type(extra)}") + def cleanup(self): + for extra in self.get_extras_list(): + extra.cleanup() + def clone(self): cloned = ContextExtrasGroup() cloned.context_ref = self.context_ref diff --git a/animatediff/nodes_context.py b/animatediff/nodes_context.py index 38c194d..8f69aad 100644 --- a/animatediff/nodes_context.py +++ b/animatediff/nodes_context.py @@ -1,5 +1,6 @@ import torch from torch import Tensor +from typing import Union import comfy.samplers from comfy.model_patcher import ModelPatcher @@ -477,11 +478,11 @@ class ContextExtras_NaiveReuse: }, "optional": { "prev_extras": ("CONTEXT_EXTRAS",), - "mask_opt": ("MASK",), + "strength_multival": ("MULTIVAL",), "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), "end_percent": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001}), "weighted_mean": ("FLOAT", {"default": 0.95, "min": 0.0, "max": 1.0, "step": 0.001}), - "autosize": ("ADEAUTOSIZE", {"padding": 55}), + "autosize": ("ADEAUTOSIZE", {"padding": 0}), } } @@ -489,12 +490,13 @@ class ContextExtras_NaiveReuse: CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras" FUNCTION = "create_context_extra" - def create_context_extra(self, start_percent=0.0, end_percent=0.1, weighted_mean=0.95, mask_opt: Tensor=None, prev_extras: ContextExtrasGroup=None): + def create_context_extra(self, start_percent=0.0, end_percent=0.1, weighted_mean=0.95, strength_multival: Union[float, Tensor]=None, + prev_extras: ContextExtrasGroup=None): if prev_extras is None: prev_extras = prev_extras = ContextExtrasGroup() prev_extras = prev_extras.clone() # create extra - naive_reuse = NaiveReuse(start_percent=start_percent, end_percent=end_percent, weighted_mean=weighted_mean, mask_opt=mask_opt) + naive_reuse = NaiveReuse(start_percent=start_percent, end_percent=end_percent, weighted_mean=weighted_mean, multival_opt=strength_multival) prev_extras.add(naive_reuse) return (prev_extras,) @@ -507,10 +509,10 @@ class ContextExtras_ContextRef: }, "optional": { "prev_extras": ("CONTEXT_EXTRAS",), - "mask_opt": ("MASK",), + "strength_multival": ("MULTIVAL",), "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": 55}), + "autosize": ("ADEAUTOSIZE", {"padding": 0}), } } @@ -518,7 +520,8 @@ class ContextExtras_ContextRef: CATEGORY = "Animate Diff 🎭🅐🅓/context opts/context extras" FUNCTION = "create_context_extra" - def create_context_extra(self, start_percent=0.0, end_percent=0.1, mask_opt: Tensor=None, prev_extras: ContextExtrasGroup=None): + def create_context_extra(self, start_percent=0.0, end_percent=0.1, strength_multival: Union[float, Tensor]=None, + prev_extras: ContextExtrasGroup=None): if prev_extras is None: prev_extras = prev_extras = ContextExtrasGroup() prev_extras = prev_extras.clone() diff --git a/animatediff/nodes_multival.py b/animatediff/nodes_multival.py index 1149104..4579e17 100644 --- a/animatediff/nodes_multival.py +++ b/animatediff/nodes_multival.py @@ -4,7 +4,7 @@ from typing import Union import torch from torch import Tensor -from .utils_motion import linear_conversion, normalize_min_max, extend_to_batch_size, extend_list_to_batch_size +from .utils_motion import create_multival_combo, linear_conversion, normalize_min_max, extend_to_batch_size, extend_list_to_batch_size class ScaleType: @@ -31,39 +31,7 @@ class MultivalDynamicNode: FUNCTION = "create_multival" def create_multival(self, float_val: Union[float, list[float]]=1.0, mask_optional: Tensor=None): - # first, normalize inputs - # if float_val is iterable, treat as a list and assume inputs are floats - float_is_iterable = False - if isinstance(float_val, Iterable): - float_is_iterable = True - float_val = list(float_val) - # if mask present, make sure float_val list can be applied to list - match lengths - if mask_optional is not None: - if len(float_val) < mask_optional.shape[0]: - # copies last entry enough times to match mask shape - float_val = extend_list_to_batch_size(float_val, mask_optional.shape[0]) - if mask_optional.shape[0] < len(float_val): - mask_optional = extend_to_batch_size(mask_optional, len(float_val)) - float_val = float_val[:mask_optional.shape[0]] - float_val: Tensor = torch.tensor(float_val).unsqueeze(-1).unsqueeze(-1) - # now that inputs are normalized, figure out what value to actually return - if mask_optional is not None: - mask_optional = mask_optional.clone() - if float_is_iterable: - mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device) - else: - mask_optional = mask_optional * float_val - return (mask_optional,) - else: - if not float_is_iterable: - return (float_val,) - # create a dummy mask of b,h,w=float_len,1,1 (sigle pixel) - # purpose is for float input to work with mask code, without special cases - float_len = float_val.shape[0] if float_is_iterable else 1 - shape = (float_len,1,1) - mask_optional = torch.ones(shape) - mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device) - return (mask_optional,) + return (create_multival_combo(float_val=float_val, mask_optional=mask_optional),) class MultivalScaledMaskNode: diff --git a/animatediff/sampling.py b/animatediff/sampling.py index ddbde56..72c1794 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -114,6 +114,7 @@ class AnimateDiffHelper_GlobalState: del self.motion_models self.motion_models = None if self.params is not None: + self.params.context_options.reset() del self.params self.params = None if self.sample_settings is not None: @@ -894,7 +895,7 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options new_ctx_idxs = [zz for zz in list(range(z, z+len(cached_naive_ctx_idxs))) if zz < ADGS.params.full_length] # make sure when getting cached_naive idxs, they are adjusted for actual length leftover length adjusted_cnaive_ctx_idxs = cached_naive_ctx_idxs[:len(new_ctx_idxs)] - weighted_mean = ADGS.params.context_options.extras.naive_reuse.weighted_mean + weighted_mean = ADGS.params.context_options.extras.naive_reuse.get_effective_weighted_mean(x_in, new_ctx_idxs) conds_final[i][new_ctx_idxs] = (weighted_mean * (cached_naive_conds[i][adjusted_cnaive_ctx_idxs]*counts_final[i][new_ctx_idxs])) + ((1.-weighted_mean) * conds_final[i][new_ctx_idxs]) #conds_final[i][new_idxs] += (cached_naive_conds[i][cached_naive_full_idxs] / cached_naive_counts[i][cached_naive_full_idxs]) * counts #counts = counts_final[i][new_idxs] * naive_counts_mult# / 2 diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index fcc205a..0cf5b1a 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -3,6 +3,7 @@ import torch import torch.nn.functional as F from torch import Tensor, nn from abc import ABC, abstractmethod +from collections.abc import Iterable import comfy.model_management as model_management import comfy.ops @@ -238,6 +239,42 @@ class InputPIA_Multival(InputPIA): return mask * self.multival +def create_multival_combo(float_val: Union[float, list[float]], mask_optional: Tensor=None): + # first, normalize inputs + # if float_val is iterable, treat as a list and assume inputs are floats + float_is_iterable = False + if isinstance(float_val, Iterable): + float_is_iterable = True + float_val = list(float_val) + # if mask present, make sure float_val list can be applied to list - match lengths + if mask_optional is not None: + if len(float_val) < mask_optional.shape[0]: + # copies last entry enough times to match mask shape + float_val = extend_list_to_batch_size(float_val, mask_optional.shape[0]) + if mask_optional.shape[0] < len(float_val): + mask_optional = extend_to_batch_size(mask_optional, len(float_val)) + float_val = float_val[:mask_optional.shape[0]] + float_val: Tensor = torch.tensor(float_val).unsqueeze(-1).unsqueeze(-1) + # now that inputs are normalized, figure out what value to actually return + if mask_optional is not None: + mask_optional = mask_optional.clone() + if float_is_iterable: + mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device) + else: + mask_optional = mask_optional * float_val + return mask_optional + else: + if not float_is_iterable: + return float_val + # create a dummy mask of b,h,w=float_len,1,1 (sigle pixel) + # purpose is for float input to work with mask code, without special cases + float_len = float_val.shape[0] if float_is_iterable else 1 + shape = (float_len,1,1) + mask_optional = torch.ones(shape) + mask_optional = mask_optional[:] * float_val.to(mask_optional.dtype).to(mask_optional.device) + return mask_optional + + def get_combined_multival(multivalA: Union[float, Tensor], multivalB: Union[float, Tensor]) -> Union[float, Tensor]: # if one is None, use the other if multivalA == None: