diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py deleted file mode 100644 index 455257a..0000000 --- a/animatediff/conditioning.py +++ /dev/null @@ -1,303 +0,0 @@ -from torch import Tensor - -from comfy.model_base import BaseModel - -from .utils_motion import get_sorted_list_via_attr - - -class LoraHookMode: - MIN_VRAM = "min_vram" - MAX_SPEED = "max_speed" - #MIN_VRAM_LOWVRAM = "min_vram_lowvram" - #MAX_SPEED_LOWVRAM = "max_speed_lowvram" - - -# Acts simply as a way to track unique LoraHooks -class HookRef: - pass - - -class LoraHook: - def __init__(self, lora_name: str): - self.lora_name = lora_name - self.lora_keyframe = LoraHookKeyframeGroup() - self.hook_ref = HookRef() - - def initialize_timesteps(self, model: BaseModel): - self.lora_keyframe.initialize_timesteps(model) - - def reset(self): - self.lora_keyframe.reset() - - - def get_copy(self): - ''' - Copies LoraHook, but maintains same HookRef - ''' - c = LoraHook(lora_name=self.lora_name) - c.lora_keyframe = self.lora_keyframe - c.hook_ref = self.hook_ref # same instance that acts as ref - return c - - @property - def strength(self): - return self.lora_keyframe.strength - - def __eq__(self, other: 'LoraHook'): - return self.__class__ == other.__class__ and self.hook_ref == other.hook_ref - - def __hash__(self): - return hash(self.hook_ref) - - -class LoraHookGroup: - ''' - Stores LoRA hooks to apply for conditioning - ''' - def __init__(self): - self.hooks: list[LoraHook] = [] - - def names(self): - names = [] - for hook in self.hooks: - names.append(hook.lora_name) - return ",".join(names) - - def add(self, hook: LoraHook): - if hook not in self.hooks: - self.hooks.append(hook) - - def is_empty(self): - return len(self.hooks) == 0 - - def contains(self, lora_hook: LoraHook): - return lora_hook in self.hooks - - def clone(self): - cloned = LoraHookGroup() - for hook in self.hooks: - cloned.add(hook.get_copy()) - return cloned - - def clone_and_combine(self, other: 'LoraHookGroup'): - cloned = self.clone() - for hook in other.hooks: - cloned.add(hook.get_copy()) - return cloned - - def set_keyframes_on_hooks(self, hook_kf: 'LoraHookKeyframeGroup'): - hook_kf = hook_kf.clone() - for hook in self.hooks: - hook.lora_keyframe = hook_kf - - @staticmethod - def combine_all_lora_hooks(lora_hooks_list: list['LoraHookGroup'], require_count=1) -> 'LoraHookGroup': - actual: list[LoraHookGroup] = [] - for group in lora_hooks_list: - if group is not None: - actual.append(group) - if len(actual) < require_count: - raise Exception(f"Need at least {require_count} LoRA Hooks to combine, but only had {len(actual)}.") - # if only 1 hook, just return itself without any cloning - if len(actual) == 1: - return actual[0] - final_hook: LoraHookGroup = None - for hook in actual: - if final_hook is None: - final_hook = hook.clone() - else: - final_hook = final_hook.clone_and_combine(hook) - return final_hook - - -class LoraHookKeyframe: - def __init__(self, strength: float, start_percent=0.0, guarantee_steps=1): - self.strength = strength - # scheduling - self.start_percent = float(start_percent) - self.start_t = 999999999.9 - self.guarantee_steps = guarantee_steps - - def clone(self): - c = LoraHookKeyframe(strength=self.strength, - start_percent=self.start_percent, guarantee_steps=self.guarantee_steps) - c.start_t = self.start_t - return c - -class LoraHookKeyframeGroup: - def __init__(self): - self.keyframes: list[LoraHookKeyframe] = [] - self._current_keyframe: LoraHookKeyframe = None - self._current_used_steps: int = 0 - self._current_index: int = 0 - self._curr_t: float = -1 - - def reset(self): - self._current_keyframe = None - self._current_used_steps = 0 - self._current_index = 0 - self._curr_t = -1 - self._set_first_as_current() - - def add(self, keyframe: LoraHookKeyframe): - # add to end of list, then sort - self.keyframes.append(keyframe) - self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent") - self._set_first_as_current() - - def _set_first_as_current(self): - if len(self.keyframes) > 0: - self._current_keyframe = self.keyframes[0] - else: - self._current_keyframe = None - - def has_index(self, index: int) -> int: - return index >= 0 and index < len(self.keyframes) - - def is_empty(self) -> bool: - return len(self.keyframes) == 0 - - def clone(self): - cloned = LoraHookKeyframeGroup() - for keyframe in self.keyframes: - cloned.keyframes.append(keyframe) - cloned._set_first_as_current() - return cloned - - def initialize_timesteps(self, model: BaseModel): - for keyframe in self.keyframes: - keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent) - - def prepare_current_keyframe(self, curr_t: float) -> bool: - if self.is_empty(): - return False - if curr_t == self._curr_t: - return False - prev_index = self._current_index - # if met guaranteed steps, look for next keyframe in case need to switch - if self._current_used_steps >= self._current_keyframe.guarantee_steps: - # if has next index, loop through and see if need t oswitch - if self.has_index(self._current_index+1): - for i in range(self._current_index+1, len(self.keyframes)): - eval_c = self.keyframes[i] - # check if start_t is greater or equal to curr_t - # NOTE: t is in terms of sigmas, not percent, so bigger number = earlier step in sampling - if eval_c.start_t >= curr_t: - self._current_index = i - self._current_keyframe = eval_c - self._current_used_steps = 0 - # if guarantee_steps greater than zero, stop searching for other keyframes - if self._current_keyframe.guarantee_steps > 0: - break - # if eval_c is outside the percent range, stop looking further - else: break - # update steps current context is used - self._current_used_steps += 1 - # update current timestep this was performed on - self._curr_t = curr_t - # return True if keyframe changed, False if no change - return prev_index != self._current_index - - # properties shadow those of LoraHookKeyframe - @property - def strength(self): - if self._current_keyframe is not None: - return self._current_keyframe.strength - return 1.0 - - -class COND_CONST: - KEY_LORA_HOOK = "lora_hook" - KEY_DEFAULT_COND = "default_cond" - - COND_AREA_DEFAULT = "default" - COND_AREA_MASK_BOUNDS = "mask bounds" - _LIST_COND_AREA = [COND_AREA_DEFAULT, COND_AREA_MASK_BOUNDS] - - -class TimestepsCond: - def __init__(self, start_percent: float, end_percent: float): - self.start_percent = start_percent - self.end_percent = end_percent - - -def conditioning_set_values(conditioning, values={}): - c = [] - for t in conditioning: - n = [t[0], t[1].copy()] - for k in values: - n[1][k] = values[k] - c.append(n) - return c - -def set_lora_hook_for_conditioning(conditioning, lora_hook: LoraHookGroup): - if lora_hook is None: - return conditioning - return conditioning_set_values(conditioning, {COND_CONST.KEY_LORA_HOOK: lora_hook}) - -def set_timesteps_for_conditioning(conditioning, timesteps_cond: TimestepsCond): - if timesteps_cond is None: - return conditioning - return conditioning_set_values(conditioning, {"start_percent": timesteps_cond.start_percent, - "end_percent": timesteps_cond.end_percent}) - -def set_mask_for_conditioning(conditioning, mask: Tensor, set_cond_area: str, strength: float): - if mask is None: - return conditioning - set_area_to_bounds = False - if set_cond_area != COND_CONST.COND_AREA_DEFAULT: - set_area_to_bounds = True - if len(mask.shape) < 3: - mask = mask.unsqueeze(0) - - return conditioning_set_values(conditioning, {"mask": mask, - "set_area_to_bounds": set_area_to_bounds, - "mask_strength": strength}) - -def combine_conditioning(conds: list): - combined_conds = [] - for cond in conds: - combined_conds.extend(cond) - return combined_conds - -def set_mask_conds(conds: list, strength: float, set_cond_area: str, - opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None): - masked_conds = [] - for c in conds: - # first, apply lora_hook to conditioning, if provided - c = set_lora_hook_for_conditioning(c, opt_lora_hook) - # next, apply mask to conditioning - c = set_mask_for_conditioning(conditioning=c, mask=opt_mask, strength=strength, set_cond_area=set_cond_area) - # apply timesteps, if present - c = set_timesteps_for_conditioning(conditioning=c, timesteps_cond=opt_timesteps) - # finally, apply mask to conditioning and store - masked_conds.append(c) - return masked_conds - -def set_mask_and_combine_conds(conds: list, new_conds: list, strength: float=1.0, set_cond_area: str="default", - opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None): - combined_conds = [] - for c, masked_c in zip(conds, new_conds): - # first, apply lora_hook to new conditioning, if provided - masked_c = set_lora_hook_for_conditioning(masked_c, opt_lora_hook) - # next, apply mask to new conditioning, if provided - masked_c = set_mask_for_conditioning(conditioning=masked_c, mask=opt_mask, set_cond_area=set_cond_area, strength=strength) - # apply timesteps, if present - masked_c = set_timesteps_for_conditioning(conditioning=masked_c, timesteps_cond=opt_timesteps) - # finally, combine with existing conditioning and store - combined_conds.append(combine_conditioning([c, masked_c])) - return combined_conds - -def set_unmasked_and_combine_conds(conds: list, new_conds: list, - opt_lora_hook: LoraHookGroup, opt_timesteps: TimestepsCond=None): - combined_conds = [] - for c, new_c in zip(conds, new_conds): - # first, apply lora_hook to new conditioning, if provided - new_c = set_lora_hook_for_conditioning(new_c, opt_lora_hook) - # next, add default_cond key to cond so that during sampling, it can be identified - new_c = conditioning_set_values(new_c, {COND_CONST.KEY_DEFAULT_COND: True}) - # apply timesteps, if present - new_c = set_timesteps_for_conditioning(conditioning=new_c, timesteps_cond=opt_timesteps) - # finally, combine with existing conditioning and store - combined_conds.append(combine_conditioning([c, new_c])) - return combined_conds diff --git a/animatediff/context.py b/animatediff/context.py index df116d1..106937d 100644 --- a/animatediff/context.py +++ b/animatediff/context.py @@ -12,6 +12,7 @@ from comfy.model_base import BaseModel from comfy.model_patcher import ModelPatcher from .context_extras import ContextExtrasGroup +from .utils_model import BIGMAX from .utils_motion import get_sorted_list_via_attr @@ -65,6 +66,12 @@ class ContextOptions: if self.view_options: self.view_options.step = value + def get_effective_guarantee_steps(self, max_sigma: torch.Tensor): + '''If keyframe starts before current sampling range (max_sigma), treat as 0.''' + if self.start_t > max_sigma: + return 0 + return self.guarantee_steps + def clone(self): n = ContextOptions(context_length=self.context_length, context_stride=self.context_stride, context_overlap=self.context_overlap, context_schedule=self.context_schedule, @@ -141,18 +148,19 @@ class ContextOptionsGroup: context.start_t = model.model_sampling.percent_to_sigma(context.start_percent) self.extras.initialize_timesteps(model) - def prepare_current(self, t: Tensor): - self.prepare_current_context(t) - self.extras.prepare_current(t) + def prepare_current(self, t: Tensor, transformer_options): + self.prepare_current_context(t, transformer_options) + self.extras.prepare_current(t, transformer_options) - def prepare_current_context(self, t: Tensor): + def prepare_current_context(self, t: Tensor, transformer_options: dict[str, Tensor]): curr_t: float = t[0] # if same as previous, do nothing as step already accounted for if curr_t == self._previous_t: return prev_index = self._current_index + max_sigma = torch.max(transformer_options.get("sigmas", BIGMAX)) # if met guaranteed steps, look for next context in case need to switch - if self._current_used_steps >= self._current_context.guarantee_steps: + if self._current_used_steps >= self._current_context.get_effective_guarantee_steps(max_sigma): # if has next index, loop through and see if need to switch if self.has_index(self._current_index+1): for i in range(self._current_index+1, len(self.contexts)): @@ -164,7 +172,7 @@ class ContextOptionsGroup: self._current_context = eval_c self._current_used_steps = 0 # if guarantee_steps greater than zero, stop searching for other keyframes - if self._current_context.guarantee_steps > 0: + if self._current_context.get_effective_guarantee_steps(max_sigma) > 0: break # if eval_c is outside the percent range, stop looking further else: diff --git a/animatediff/context_extras.py b/animatediff/context_extras.py index 2f19856..dc7d959 100644 --- a/animatediff/context_extras.py +++ b/animatediff/context_extras.py @@ -5,6 +5,7 @@ from torch import Tensor from comfy.model_base import BaseModel +from .utils_model import BIGMAX from .utils_motion import (prepare_mask_batch, extend_to_batch_size, get_combined_multival, resize_multival, get_sorted_list_via_attr) @@ -25,7 +26,7 @@ class ContextExtra: self.start_t = model.model_sampling.percent_to_sigma(self.start_percent) self.end_t = model.model_sampling.percent_to_sigma(self.end_percent) - def prepare_current(self, t: Tensor): + def prepare_current(self, t: Tensor, transformer_options: dict[str, Tensor]): self.curr_t = t[0] def should_run(self): @@ -260,6 +261,12 @@ class NaiveReuseKeyframe: self.guarantee_steps = guarantee_steps self.inherit_missing = inherit_missing + def get_effective_guarantee_steps(self, max_sigma: torch.Tensor): + '''If keyframe starts before current sampling range (max_sigma), treat as 0.''' + if self.start_t > max_sigma: + return 0 + return self.guarantee_steps + def clone(self): c = NaiveReuseKeyframe(mult=self.mult, mult_multival=self.mult_multival, start_percent=self.start_percent, guarantee_steps=self.guarantee_steps) @@ -330,7 +337,7 @@ class NaiveReuseKeyframeGroup: for keyframe in self.keyframes: keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent) - def prepare_current_keyframe(self, t: Tensor): + def prepare_current_keyframe(self, t: Tensor, transformer_options: dict[str, Tensor]): if self.is_empty(): return curr_t: float = t[0] @@ -338,8 +345,9 @@ class NaiveReuseKeyframeGroup: if curr_t == self._previous_t: return prev_index = self._current_index + max_sigma = torch.max(transformer_options.get("sigmas", BIGMAX)) # if met guaranteed steps, look for next keyframe in case need to switch - if self._current_used_steps >= self._current_keyframe.guarantee_steps: + if self._current_used_steps >= self._current_keyframe.get_effective_guarantee_steps(max_sigma): # if has next index, loop through and see if need t oswitch if self.has_index(self._current_index+1): for i in range(self._current_index+1, len(self.keyframes)): @@ -351,7 +359,7 @@ class NaiveReuseKeyframeGroup: self._current_keyframe = eval_c self._current_used_steps = 0 # if guarantee_steps greater than zero, stop searching for other keyframes - if self._current_keyframe.guarantee_steps > 0: + if self._current_keyframe.get_effective_guarantee_steps(max_sigma) > 0: break # if eval_c is outside the percent range, stop looking further else: break @@ -394,9 +402,9 @@ class NaiveReuse(ContextExtra): super().initialize_timesteps(model) self.keyframe.initialize_timesteps(model) - def prepare_current(self, t: Tensor): - super().prepare_current(t) - self.keyframe.prepare_current_keyframe(t) + def prepare_current(self, t: Tensor, transformer_options: dict[str, Tensor]): + super().prepare_current(t, transformer_options) + self.keyframe.prepare_current_keyframe(t, transformer_options) def get_effective_weighted_mean(self, x: Tensor, idxs: list[int]): if self.orig_multival is None and self.keyframe.mult_multival is None: @@ -427,6 +435,10 @@ class NaiveReuse(ContextExtra): #-------------------------------- +################################ +# DenoiseReuse + + class ContextExtrasGroup: def __init__(self): self.context_ref: ContextRef = None @@ -444,9 +456,9 @@ class ContextExtrasGroup: for extra in self.get_extras_list(): extra.initialize_timesteps(model) - def prepare_current(self, t: Tensor): + def prepare_current(self, t: Tensor, transformer_options): for extra in self.get_extras_list(): - extra.prepare_current(t) + extra.prepare_current(t, transformer_options) def should_run_context_ref(self): if not self.context_ref: diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 0a4aca7..6ce1144 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -29,9 +29,8 @@ from .utils_motion import (ADKeyframe, ADKeyframeGroup, MotionCompatibilityError PerBlock, AllPerBlocks, get_combined_per_block_list, get_combined_multival, get_combined_input, get_combined_input_effect_multival, ade_broadcast_image_to, extend_to_batch_size, prepare_mask_batch) -from .conditioning import HookRef, LoraHook, LoraHookGroup, LoraHookMode from .motion_lora import MotionLoraInfo, MotionLoraList -from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type, vae_encode_raw_batched +from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type, vae_encode_raw_batched, BIGMAX from .sample_settings import SampleSettings, SeedNoiseGeneration from .dinklink import DinkLinkConst, get_dinklink, get_acn_outer_sample_wrapper @@ -328,14 +327,15 @@ class MotionModelAttachment: for keyframe in self.keyframes.keyframes: keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent) - def prepare_current_keyframe(self, patcher: MotionModelPatcher, x: Tensor, t: Tensor): + def prepare_current_keyframe(self, patcher: MotionModelPatcher, x: Tensor, t: Tensor, transformer_options: dict[str, Tensor]): curr_t: float = t[0] # if curr_t was previous_t, then do nothing (already accounted for this step) if curr_t == self.previous_t: return prev_index = self.current_index + max_sigma = torch.max(transformer_options.get("sigmas", BIGMAX)) # if met guaranteed steps, look for next keyframe in case need to switch - if self.current_keyframe is None or self.current_used_steps >= self.current_keyframe.guarantee_steps: + if self.current_keyframe is None or self.current_used_steps >= self.current_keyframe.get_effective_guarantee_steps(max_sigma): # if has next index, loop through and see if need to switch if self.keyframes.has_index(self.current_index+1): for i in range(self.current_index+1, len(self.keyframes)): @@ -373,7 +373,7 @@ class MotionModelAttachment: elif not self.current_keyframe.inherit_missing: self.current_pia_input = None # if guarantee_steps greater than zero, stop searching for other keyframes - if self.current_keyframe.guarantee_steps > 0: + if self.current_keyframe.get_effective_guarantee_steps(max_sigma) > 0: break # if eval_kf is outside the percent range, stop looking further else: @@ -723,10 +723,10 @@ class MotionModelGroup: for motion_model in self.models: motion_model.cleanup() - def prepare_current_keyframe(self, x: Tensor, t: Tensor): + def prepare_current_keyframe(self, x: Tensor, t: Tensor, transformer_options: dict[str, Tensor]): for motion_model in self.models: attachment = get_mm_attachment(motion_model) - attachment.prepare_current_keyframe(motion_model, x=x, t=t) + attachment.prepare_current_keyframe(motion_model, x=x, t=t, transformer_options=transformer_options) def get_special_models(self): pia_motion_models: list[MotionModelPatcher] = [] diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 54af3e5..7bd09f6 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -12,7 +12,6 @@ import comfy_extras.nodes_hooks import comfy.hooks import comfy.utils -from .conditioning import (COND_CONST) from .utils_model import BIGMAX, InterpolationMethod from .logger import logger @@ -25,6 +24,12 @@ from .logger import logger #------------------------------------------------------------------ #------------------------------------------------------------------ #------------------------------------------------------------------ +class COND_CONST: + COND_AREA_DEFAULT = "default" + COND_AREA_MASK_BOUNDS = "mask bounds" + _LIST_COND_AREA = [COND_AREA_DEFAULT, COND_AREA_MASK_BOUNDS] + + class CreateLoraHookKeyframeInterpolationDEPR: @classmethod def INPUT_TYPES(s): diff --git a/animatediff/sample_settings.py b/animatediff/sample_settings.py index b7eeb38..d504875 100644 --- a/animatediff/sample_settings.py +++ b/animatediff/sample_settings.py @@ -13,9 +13,8 @@ from comfy.model_base import BaseModel from comfy.sd import VAE from . import freeinit -from .conditioning import LoraHookMode from .context import ContextOptions, ContextOptionsGroup -from .utils_model import SigmaSchedule +from .utils_model import SigmaSchedule, BIGMAX from .utils_motion import extend_to_batch_size, get_sorted_list_via_attr, prepare_mask_batch from .logger import logger @@ -611,6 +610,12 @@ class CustomCFGKeyframe: self.start_t = 999999999.9 self.guarantee_steps = guarantee_steps + def get_effective_guarantee_steps(self, max_sigma: torch.Tensor): + '''If keyframe starts before current sampling range (max_sigma), treat as 0.''' + if self.start_t > max_sigma: + return 0 + return self.guarantee_steps + def clone(self): c = CustomCFGKeyframe(cfg_multival=self.cfg_multival, start_percent=self.start_percent, guarantee_steps=self.guarantee_steps) @@ -661,14 +666,15 @@ class CustomCFGKeyframeGroup: for keyframe in self.keyframes: keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent) - def prepare_current_keyframe(self, t: Tensor): + def prepare_current_keyframe(self, t: Tensor, transformer_options: dict[str, Tensor]): curr_t: float = t[0] # if curr_t same as before, do nothing as step already accounted for if curr_t == self._previous_t: return prev_index = self._current_index + max_sigma = torch.max(transformer_options.get("sigmas", BIGMAX)) # if met guaranteed steps, look for next keyframe in case need to switch - if self._current_used_steps >= self._current_keyframe.guarantee_steps: + if self._current_used_steps >= self._current_keyframe.get_effective_guarantee_steps(max_sigma): # if has next index, loop through and see if need t oswitch if self.has_index(self._current_index+1): for i in range(self._current_index+1, len(self.keyframes)): @@ -680,7 +686,7 @@ class CustomCFGKeyframeGroup: self._current_keyframe = eval_c self._current_used_steps = 0 # if guarantee_steps greater than zero, stop searching for other keyframes - if self._current_keyframe.guarantee_steps > 0: + if self._current_keyframe.get_effective_guarantee_steps(max_sigma) > 0: break # if eval_c is outside the percent range, stop looking further else: break diff --git a/animatediff/sampling.py b/animatediff/sampling.py index a48af04..7969ce2 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -54,13 +54,13 @@ class AnimateDiffGlobalState: if self.sample_settings.custom_cfg is not None: self.sample_settings.custom_cfg.initialize_timesteps(model) - def prepare_current_keyframes(self, x: Tensor, timestep: Tensor): + def prepare_current_keyframes(self, x: Tensor, timestep: Tensor, transformer_options: dict[str, Tensor]): if self.motion_models is not None: - self.motion_models.prepare_current_keyframe(x=x, t=timestep) + self.motion_models.prepare_current_keyframe(x=x, t=timestep, transformer_options=transformer_options) if self.params.context_options is not None: - self.params.context_options.prepare_current(t=timestep) + self.params.context_options.prepare_current(t=timestep, transformer_options=transformer_options) if self.sample_settings.custom_cfg is not None: - self.sample_settings.custom_cfg.prepare_current_keyframe(t=timestep) + self.sample_settings.custom_cfg.prepare_current_keyframe(t=timestep, transformer_options=transformer_options) def perform_special_model_features(self, model: BaseModel, conds: list, x_in: Tensor, model_options: dict[str]): if self.motion_models is not None: @@ -540,7 +540,7 @@ def outer_sample_wrapper(executor: WrapperExecutor, *args, **kwargs): def evolved_sampling_function(model, x: Tensor, timestep: Tensor, uncond, cond, cond_scale, model_options: dict={}, seed=None): ADGS: AnimateDiffGlobalState = model_options["transformer_options"]["ADGS"] ADGS.initialize(model) - ADGS.prepare_current_keyframes(x=x, timestep=timestep) + ADGS.prepare_current_keyframes(x=x, timestep=timestep, transformer_options=model_options["transformer_options"]) try: # add AD/evolved-sampling params to model_options (transformer_options) model_options = model_options.copy() diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index 122244f..362c56b 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -445,6 +445,12 @@ class ADKeyframe: def has_pia_input(self): return self.pia_input is not None + def get_effective_guarantee_steps(self, max_sigma: torch.Tensor): + '''If keyframe starts before current sampling range (max_sigma), treat as 0.''' + if self.start_t > max_sigma: + return 0 + return self.guarantee_steps + class ADKeyframeGroup: def __init__(self):