diff --git a/animatediff/context.py b/animatediff/context.py index 99ba794..60b7fd7 100644 --- a/animatediff/context.py +++ b/animatediff/context.py @@ -7,8 +7,9 @@ import numpy as np class ContextSchedules: UNIFORM = "uniform" UNIFORM_CONSTANT = "uniform_constant" + UNIFORM_V2 = "uniform v2" - CONTEXT_SCHEDULE_LIST = [UNIFORM] + CONTEXT_SCHEDULE_LIST = [UNIFORM, UNIFORM_V2] # Returns fraction that has denominator that is a power of 2 @@ -53,6 +54,51 @@ def uniform( ): yield [e % num_frames for e in range(j, j + context_size * context_step, context_step)] +def uniform_v2( + step: int = ..., + num_steps: Optional[int] = None, + num_frames: int = ..., + context_size: Optional[int] = None, + context_stride: int = 3, + context_overlap: int = 4, + closed_loop: bool = True, + print_final: bool = False, +): + if num_frames <= context_size: + yield list(range(num_frames)) + return + + context_stride = min(context_stride, int(np.ceil(np.log2(num_frames / context_size))) + 1) + + pad = int(round(num_frames * ordered_halving(step, print_final))) + for context_step in 1 << np.arange(context_stride): + j_initial = int(ordered_halving(step) * context_step) + pad + for j in range( + j_initial, + num_frames + pad - context_overlap, + (context_size * context_step - context_overlap), + ): + if context_size * context_step > num_frames: + # On the final context_step, + # ensure no frame appears in the window twice + yield [e % num_frames for e in range(j, j + num_frames, context_step)] + continue + j = j % num_frames + if j > (j + context_size * context_step) % num_frames and not closed_loop: + yield [e for e in range(j, num_frames, context_step)] + j_stop = (j + context_size * context_step) % num_frames + # When ((num_frames % (context_size - context_overlap)+context_overlap) % context_size != 0, + # This can cause 'superflous' runs where all frames in + # a context window have already been processed during + # the first context window of this stride and step. + # While the following commented if should prevent this, + # I believe leaving it in is more correct as it maintains + # the total conditional passes per frame over a large total steps + # if j_stop > context_overlap: + yield [e for e in range(0, j_stop, context_step)] + continue + yield [e % num_frames for e in range(j, j + context_size * context_step, context_step)] + def uniform_constant( step: int = ..., @@ -103,6 +149,8 @@ def get_context_scheduler(name: str) -> Callable: return uniform case ContextSchedules.UNIFORM_CONSTANT: return uniform_constant + case ContextSchedules.UNIFORM_V2: + return uniform_v2 case _: raise ValueError(f"Unknown context_overlap policy {name}")