Views, start of Context scheduling, use_on_equal_length for Contexts,
This commit is contained in:
+1
-1
@@ -1,6 +1,6 @@
|
||||
import folder_paths
|
||||
from .animatediff.logger import logger
|
||||
from .animatediff.model_utils import get_available_motion_models, Folders
|
||||
from .animatediff.utils_model import get_available_motion_models, Folders
|
||||
from .animatediff.nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
if len(get_available_motion_models()) == 0:
|
||||
|
||||
+157
-183
@@ -1,7 +1,9 @@
|
||||
from typing import Callable, Optional
|
||||
from typing import Callable, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .utils_motion import get_sorted_list_via_attr
|
||||
|
||||
class ContextFuseMethod:
|
||||
FLAT = "flat"
|
||||
PYRAMID = "pyramid"
|
||||
@@ -15,22 +17,128 @@ class ContextType:
|
||||
|
||||
class ContextOptions:
|
||||
def __init__(self, context_length: int=None, context_stride: int=None, context_overlap: int=None,
|
||||
context_schedule: str=None, closed_loop: bool=False, fuse_method: str=ContextFuseMethod.FLAT):
|
||||
context_schedule: str=None, closed_loop: bool=False, fuse_method: str=ContextFuseMethod.FLAT,
|
||||
use_on_equal_length: bool=False, view_options: 'ContextOptions'=None,
|
||||
start_percent=0.0, guarantee_steps=1):
|
||||
# permanent settings
|
||||
self.context_length = context_length
|
||||
self.context_stride = context_stride
|
||||
self.context_overlap = context_overlap
|
||||
self.context_schedule = context_schedule
|
||||
self.closed_loop = closed_loop
|
||||
self.fuse_method = fuse_method
|
||||
self.sync_context_to_pe = False
|
||||
self.sync_context_to_pe = False # this feature is likely bad and stay unused, so I might remove this
|
||||
self.use_on_equal_length = use_on_equal_length
|
||||
self.view_options = view_options.clone() if view_options else view_options
|
||||
# scheduling
|
||||
self.start_percent = float(start_percent)
|
||||
self.guarantee_steps = guarantee_steps
|
||||
# temporary vars
|
||||
self._step: int = 0
|
||||
|
||||
@property
|
||||
def step(self):
|
||||
return self._step
|
||||
@step.setter
|
||||
def step(self, value: int):
|
||||
self._step = value
|
||||
if self.view_options:
|
||||
self.view_options.step = value
|
||||
|
||||
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,
|
||||
closed_loop=self.closed_loop, fuse_method=self.fuse_method)
|
||||
closed_loop=self.closed_loop, fuse_method=self.fuse_method,
|
||||
use_on_equal_length=self.use_on_equal_length, view_options=self.view_options,
|
||||
start_percent=self.start_percent, guarantee_steps=self.guarantee_steps)
|
||||
return n
|
||||
|
||||
|
||||
class ContextOptionsGroup:
|
||||
def __init__(self):
|
||||
self.contexts: list[ContextOptions] = []
|
||||
self._current_context: ContextOptions = None
|
||||
self._current_guaranteed_steps: int = 0
|
||||
self.step = 0
|
||||
|
||||
def reset(self):
|
||||
self._current_context: ContextOptions = None
|
||||
self._current_guaranteed_steps: int = 0
|
||||
self.step = 0
|
||||
self._set_first_as_current()
|
||||
|
||||
@classmethod
|
||||
def default(cls):
|
||||
def_context = ContextOptions()
|
||||
new_group = ContextOptionsGroup()
|
||||
new_group.add(def_context)
|
||||
return new_group
|
||||
|
||||
def add(self, context: ContextOptions):
|
||||
# add to end of list, then sort
|
||||
self.contexts.append(context)
|
||||
self.contexts = get_sorted_list_via_attr(self.contexts, "start_percent")
|
||||
self._set_first_as_current()
|
||||
|
||||
def add_to_start(self, context: ContextOptions):
|
||||
# add to start of list, then sort
|
||||
self.contexts.insert(0, context)
|
||||
self.contexts = get_sorted_list_via_attr(self.contexts, "start_percent")
|
||||
self._set_first_as_current()
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return len(self.contexts) == 0
|
||||
|
||||
def clone(self):
|
||||
cloned = ContextOptionsGroup()
|
||||
for context in self.contexts:
|
||||
cloned.contexts.append(context)
|
||||
cloned._set_first_as_current()
|
||||
return cloned
|
||||
|
||||
def update_current_context(self, t: float):
|
||||
self._current_context = self.contexts[0]
|
||||
# based on t + current_steps, determine which context to use
|
||||
pass
|
||||
|
||||
def _set_first_as_current(self):
|
||||
if len(self.contexts) > 0:
|
||||
self._current_context = self.contexts[0]
|
||||
|
||||
# properties shadow those of ContextOptions
|
||||
@property
|
||||
def context_length(self):
|
||||
return self._current_context.context_length
|
||||
|
||||
@property
|
||||
def context_overlap(self):
|
||||
return self._current_context.context_overlap
|
||||
|
||||
@property
|
||||
def context_stride(self):
|
||||
return self._current_context.context_stride
|
||||
|
||||
@property
|
||||
def context_schedule(self):
|
||||
return self._current_context.context_schedule
|
||||
|
||||
@property
|
||||
def closed_loop(self):
|
||||
return self._current_context.closed_loop
|
||||
|
||||
@property
|
||||
def fuse_method(self):
|
||||
return self._current_context.fuse_method
|
||||
|
||||
@property
|
||||
def use_on_equal_length(self):
|
||||
return self._current_context.use_on_equal_length
|
||||
|
||||
@property
|
||||
def view_options(self):
|
||||
return self._current_context.view_options
|
||||
|
||||
|
||||
class ContextSchedules:
|
||||
UNIFORM_LOOPED = "uniform"
|
||||
UNIFORM_STANDARD = "uniform_standard"
|
||||
@@ -39,23 +147,25 @@ class ContextSchedules:
|
||||
|
||||
BATCHED = "batched"
|
||||
|
||||
UNIFORM_SCHEDULE_LIST = [UNIFORM_LOOPED] # only include somewhat functional contexts here
|
||||
VIEW_AS_CONTEXT = "view_as_context"
|
||||
|
||||
UNIFORM_SCHEDULE_LIST = [UNIFORM_LOOPED]
|
||||
STATIC_SCHEDULE_LIST = [STATIC_STANDARD]
|
||||
|
||||
|
||||
# from https://github.com/neggles/animatediff-cli/blob/main/src/animatediff/pipelines/context.py
|
||||
def create_windows_uniform_looped(step: int, num_frames: int, opts: ContextOptions):
|
||||
def create_windows_uniform_looped(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]):
|
||||
windows = []
|
||||
if num_frames <= opts.context_length:
|
||||
if num_frames < opts.context_length:
|
||||
windows.append(list(range(num_frames)))
|
||||
return windows
|
||||
|
||||
context_stride = min(opts.context_stride, int(np.ceil(np.log2(num_frames / opts.context_length))) + 1)
|
||||
# obtain uniform windows as normal, looping and all
|
||||
for context_step in 1 << np.arange(context_stride):
|
||||
pad = int(round(num_frames * ordered_halving(step)))
|
||||
pad = int(round(num_frames * ordered_halving(opts.step)))
|
||||
for j in range(
|
||||
int(ordered_halving(step) * context_step) + pad,
|
||||
int(ordered_halving(opts.step) * context_step) + pad,
|
||||
num_frames + pad + (0 if opts.closed_loop else -opts.context_overlap),
|
||||
(opts.context_length * context_step - opts.context_overlap),
|
||||
):
|
||||
@@ -64,7 +174,7 @@ def create_windows_uniform_looped(step: int, num_frames: int, opts: ContextOptio
|
||||
return windows
|
||||
|
||||
|
||||
def create_windows_uniform_standard(step: int, num_frames: int, opts: ContextOptions):
|
||||
def create_windows_uniform_standard(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]):
|
||||
# unlike looped, uniform_straight does NOT allow windows that loop back to the beginning;
|
||||
# instead, they get shifted to the corresponding end of the frames.
|
||||
# in the case that a window (shifted or not) is identical to the previous one, it gets skipped.
|
||||
@@ -76,9 +186,9 @@ def create_windows_uniform_standard(step: int, num_frames: int, opts: ContextOpt
|
||||
context_stride = min(opts.context_stride, int(np.ceil(np.log2(num_frames / opts.context_length))) + 1)
|
||||
# first, obtain uniform windows as normal, looping and all
|
||||
for context_step in 1 << np.arange(context_stride):
|
||||
pad = int(round(num_frames * ordered_halving(step)))
|
||||
pad = int(round(num_frames * ordered_halving(opts.step)))
|
||||
for j in range(
|
||||
int(ordered_halving(step) * context_step) + pad,
|
||||
int(ordered_halving(opts.step) * context_step) + pad,
|
||||
num_frames + pad + (-opts.context_overlap),
|
||||
(opts.context_length * context_step - opts.context_overlap),
|
||||
):
|
||||
@@ -104,7 +214,7 @@ def create_windows_uniform_standard(step: int, num_frames: int, opts: ContextOpt
|
||||
break
|
||||
win_i += 1
|
||||
|
||||
# reverse delete_idxs so that they will be deleted in an order that does break idx correlation
|
||||
# reverse delete_idxs so that they will be deleted in an order that doesn't break idx correlation
|
||||
delete_idxs.reverse()
|
||||
for i in delete_idxs:
|
||||
windows.pop(i)
|
||||
@@ -112,7 +222,7 @@ def create_windows_uniform_standard(step: int, num_frames: int, opts: ContextOpt
|
||||
return windows
|
||||
|
||||
|
||||
def create_windows_static_standard(step: int, num_frames: int, opts: ContextOptions):
|
||||
def create_windows_static_standard(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]):
|
||||
windows = []
|
||||
if num_frames <= opts.context_length:
|
||||
windows.append(list(range(num_frames)))
|
||||
@@ -131,7 +241,7 @@ def create_windows_static_standard(step: int, num_frames: int, opts: ContextOpti
|
||||
return windows
|
||||
|
||||
|
||||
def create_windows_batched(step: int, num_frames: int, opts: ContextOptions):
|
||||
def create_windows_batched(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]):
|
||||
windows = []
|
||||
if num_frames <= opts.context_length:
|
||||
windows.append(list(range(num_frames)))
|
||||
@@ -144,11 +254,15 @@ def create_windows_batched(step: int, num_frames: int, opts: ContextOptions):
|
||||
return windows
|
||||
|
||||
|
||||
def get_context_windows(step: int, num_frames: int, opts: ContextOptions):
|
||||
def create_windows_default(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]):
|
||||
return [list(range(num_frames))]
|
||||
|
||||
|
||||
def get_context_windows(num_frames: int, opts: Union[ContextOptionsGroup, ContextOptions]):
|
||||
context_func = CONTEXT_MAPPING.get(opts.context_schedule, None)
|
||||
if not context_func:
|
||||
raise ValueError(f"Unknown context_schedule '{opts.context_schedule}'")
|
||||
return context_func(step, num_frames, opts)
|
||||
raise ValueError(f"Unknown context_schedule '{opts.context_schedule}'.")
|
||||
return context_func(num_frames, opts)
|
||||
|
||||
|
||||
CONTEXT_MAPPING = {
|
||||
@@ -156,19 +270,40 @@ CONTEXT_MAPPING = {
|
||||
ContextSchedules.UNIFORM_STANDARD: create_windows_uniform_standard,
|
||||
ContextSchedules.STATIC_STANDARD: create_windows_static_standard,
|
||||
ContextSchedules.BATCHED: create_windows_batched,
|
||||
ContextSchedules.VIEW_AS_CONTEXT: create_windows_default, # just return all to allow Views to do all the work
|
||||
}
|
||||
|
||||
|
||||
def generate_distance_weight(n):
|
||||
if n % 2 == 0:
|
||||
max_weight = n // 2
|
||||
def get_context_weights(num_frames: int, fuse_method: str):
|
||||
weights_func = FUSE_MAPPING.get(fuse_method, None)
|
||||
if not weights_func:
|
||||
raise ValueError(f"Unknown fuse_method '{fuse_method}'.")
|
||||
return weights_func(num_frames)
|
||||
|
||||
|
||||
def create_weights_flat(length: int, **kwargs) -> list[float]:
|
||||
# weight is the same for all
|
||||
return [1.0] * length
|
||||
|
||||
|
||||
def create_weights_pyramid(length: int, **kwargs) -> list[float]:
|
||||
# weight is based on the distance away from the edge of the context window;
|
||||
# based on weighted average concept in FreeNoise paper
|
||||
if length % 2 == 0:
|
||||
max_weight = length // 2
|
||||
weight_sequence = list(range(1, max_weight + 1, 1)) + list(range(max_weight, 0, -1))
|
||||
else:
|
||||
max_weight = (n + 1) // 2
|
||||
max_weight = (length + 1) // 2
|
||||
weight_sequence = list(range(1, max_weight, 1)) + [max_weight] + list(range(max_weight - 1, 0, -1))
|
||||
return weight_sequence
|
||||
|
||||
|
||||
FUSE_MAPPING = {
|
||||
ContextFuseMethod.FLAT: create_weights_flat,
|
||||
ContextFuseMethod.PYRAMID: create_weights_pyramid,
|
||||
}
|
||||
|
||||
|
||||
# Returns fraction that has denominator that is a power of 2
|
||||
def ordered_halving(val):
|
||||
# get binary value, padded with 0s for 64 bits
|
||||
@@ -219,164 +354,3 @@ def shift_window_to_end(window: list[int], num_frames: int):
|
||||
for i in range(len(window)):
|
||||
# 2) add end_delta to each val to slide windows to end
|
||||
window[i] = window[i] + end_delta
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
################################################################################################
|
||||
|
||||
# Generator that returns lists of latent indeces to diffuse on
|
||||
def uniform(
|
||||
step: int,
|
||||
num_frames: int,
|
||||
opts: ContextOptions,
|
||||
print_final: bool = False,
|
||||
):
|
||||
if num_frames <= opts.context_length:
|
||||
yield list(range(num_frames))
|
||||
return
|
||||
|
||||
context_stride = min(opts.context_stride, int(np.ceil(np.log2(num_frames / opts.context_length))) + 1)
|
||||
|
||||
for context_step in 1 << np.arange(context_stride):
|
||||
pad = int(round(num_frames * ordered_halving(step, print_final)))
|
||||
for j in range(
|
||||
int(ordered_halving(step) * context_step) + pad,
|
||||
num_frames + pad + (0 if opts.closed_loop else -opts.context_overlap),
|
||||
(opts.context_length * context_step - opts.context_overlap),
|
||||
):
|
||||
yield [e % num_frames for e in range(j, j + opts.context_length * context_step, context_step)]
|
||||
|
||||
|
||||
#################################
|
||||
# helper funcs for testing
|
||||
def get_total_steps(
|
||||
scheduler,
|
||||
timesteps: list[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,
|
||||
):
|
||||
return sum(
|
||||
len(
|
||||
list(
|
||||
scheduler(
|
||||
i,
|
||||
num_steps,
|
||||
num_frames,
|
||||
context_size,
|
||||
context_stride,
|
||||
context_overlap,
|
||||
)
|
||||
)
|
||||
)
|
||||
for i in range(len(timesteps))
|
||||
)
|
||||
|
||||
|
||||
def get_total_steps_fixed(
|
||||
scheduler,
|
||||
timesteps: list[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,
|
||||
):
|
||||
total_loops = 0
|
||||
for i, t in enumerate(timesteps):
|
||||
for context in scheduler(i, num_steps, num_frames, context_size, context_stride, context_overlap, closed_loop=closed_loop):
|
||||
total_loops += 1
|
||||
return total_loops
|
||||
|
||||
|
||||
def uniform_v2(
|
||||
step: int = ...,
|
||||
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 = ...,
|
||||
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)
|
||||
|
||||
# want to avoid loops that connect end to beginning
|
||||
|
||||
for context_step in 1 << np.arange(context_stride):
|
||||
pad = int(round(num_frames * ordered_halving(step, print_final)))
|
||||
for j in range(
|
||||
int(ordered_halving(step) * context_step) + pad,
|
||||
num_frames + pad + (0 if closed_loop else -context_overlap),
|
||||
(context_size * context_step - context_overlap),
|
||||
):
|
||||
skip_this_window = False
|
||||
prev_val = -1
|
||||
to_yield = []
|
||||
for e in range(j, j + context_size * context_step, context_step):
|
||||
e = e % num_frames
|
||||
# if not a closed loop and loops back on itself, should be skipped
|
||||
if not closed_loop and e < prev_val:
|
||||
skip_this_window = True
|
||||
break
|
||||
to_yield.append(e)
|
||||
prev_val = e
|
||||
if skip_this_window:
|
||||
continue
|
||||
# yield if not skipped
|
||||
yield to_yield
|
||||
|
||||
@@ -11,12 +11,12 @@ import comfy.utils
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.model_base import BaseModel
|
||||
|
||||
from .context import ContextOptions, ContextOptions
|
||||
from .context import ContextOptions, ContextOptions, ContextOptionsGroup
|
||||
from .motion_module_ad import AnimateDiffModel, has_mid_block, normalize_ad_state_dict
|
||||
from .logger import logger
|
||||
from .motion_utils import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, normalize_min_max
|
||||
from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, normalize_min_max
|
||||
from .motion_lora import MotionLoraInfo, MotionLoraList
|
||||
from .model_utils import get_motion_lora_path, get_motion_model_path, get_sd_model_type
|
||||
from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type
|
||||
from .sample_settings import SampleSettings, SeedNoiseGeneration
|
||||
|
||||
|
||||
@@ -185,7 +185,6 @@ class MotionModelPatcher(ModelPatcher):
|
||||
self.model.set_effect(self.combined_effect)
|
||||
self.was_within_range = True
|
||||
|
||||
|
||||
def cleanup(self):
|
||||
if self.model is not None:
|
||||
self.model.cleanup()
|
||||
@@ -245,6 +244,10 @@ class MotionModelGroup:
|
||||
for motion_model in self.models:
|
||||
motion_model.model.set_sub_idxs(sub_idxs=sub_idxs)
|
||||
|
||||
def set_view_options(self, view_options: ContextOptions):
|
||||
for motion_model in self.models:
|
||||
motion_model.model.set_view_options(view_options)
|
||||
|
||||
def set_video_length(self, video_length: int, full_length: int):
|
||||
for motion_model in self.models:
|
||||
motion_model.model.set_video_length(video_length=video_length, full_length=full_length)
|
||||
@@ -500,15 +503,15 @@ class InjectionParams:
|
||||
self.apply_mm_groupnorm_hack = apply_mm_groupnorm_hack
|
||||
self.model_name = model_name
|
||||
self.apply_v2_properly = apply_v2_properly
|
||||
self.context_options: ContextOptions = ContextOptions()
|
||||
self.context_options: ContextOptionsGroup = ContextOptionsGroup.default()
|
||||
self.motion_model_settings = MotionModelSettings() # Gen1
|
||||
self.sub_idxs = None # value should NOT be included in clone, so it will auto reset
|
||||
|
||||
def set_noise_extra_args(self, noise_extra_args: dict):
|
||||
noise_extra_args["context_options"] = self.context_options.clone()
|
||||
|
||||
def set_context(self, context_options: ContextOptions):
|
||||
self.context_options = context_options.clone() if context_options else ContextOptions()
|
||||
def set_context(self, context_options: ContextOptionsGroup):
|
||||
self.context_options = context_options.clone() if context_options else ContextOptionsGroup.default()
|
||||
|
||||
def is_using_sliding_context(self) -> bool:
|
||||
return self.context_options.context_length is not None
|
||||
@@ -520,7 +523,7 @@ class InjectionParams:
|
||||
self.motion_model_settings = motion_model_settings
|
||||
|
||||
def reset_context(self):
|
||||
self.context_options = ContextOptions()
|
||||
self.context_options = ContextOptionsGroup.default()
|
||||
|
||||
def clone(self) -> 'InjectionParams':
|
||||
new_params = InjectionParams(
|
||||
|
||||
@@ -13,8 +13,9 @@ from comfy.ldm.modules.diffusionmodules.openaimodel import SpatialTransformer
|
||||
from comfy.controlnet import broadcast_image_to
|
||||
from comfy.utils import repeat_to_batch_size
|
||||
|
||||
from .motion_utils import GroupNormAD, CrossAttentionMM, MotionCompatibilityError, extend_to_batch_size, prepare_mask_batch
|
||||
from .model_utils import ModelTypeSD
|
||||
from .context import ContextOptions, get_context_weights, get_context_windows
|
||||
from .utils_motion import GroupNormAD, CrossAttentionMM, MotionCompatibilityError, extend_to_batch_size, prepare_mask_batch
|
||||
from .utils_model import ModelTypeSD
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -281,6 +282,14 @@ class AnimateDiffModel(nn.Module):
|
||||
if self.mid_block is not None:
|
||||
self.mid_block.set_sub_idxs(sub_idxs)
|
||||
|
||||
def set_view_options(self, view_options: ContextOptions):
|
||||
for block in self.down_blocks:
|
||||
block.set_view_options(view_options)
|
||||
for block in self.up_blocks:
|
||||
block.set_view_options(view_options)
|
||||
if self.mid_block is not None:
|
||||
self.mid_block.set_view_options(view_options)
|
||||
|
||||
def reset(self):
|
||||
self._reset_sub_idxs()
|
||||
self._reset_scale_multiplier()
|
||||
@@ -355,6 +364,10 @@ class MotionModule(nn.Module):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.set_sub_idxs(sub_idxs)
|
||||
|
||||
def set_view_options(self, view_options: ContextOptions):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.set_view_options(view_options=view_options)
|
||||
|
||||
def reset_temp_vars(self):
|
||||
for motion_module in self.motion_modules:
|
||||
motion_module.reset_temp_vars()
|
||||
@@ -382,6 +395,7 @@ class VanillaTemporalModule(nn.Module):
|
||||
self.video_length = 16
|
||||
self.full_length = 16
|
||||
self.sub_idxs = None
|
||||
self.view_options = None
|
||||
|
||||
self.effect = None
|
||||
self.temp_effect_mask: Tensor = None
|
||||
@@ -429,8 +443,12 @@ class VanillaTemporalModule(nn.Module):
|
||||
self.sub_idxs = sub_idxs
|
||||
self.temporal_transformer.set_sub_idxs(sub_idxs)
|
||||
|
||||
def set_view_options(self, view_options: ContextOptions):
|
||||
self.view_options = view_options
|
||||
|
||||
def reset_temp_vars(self):
|
||||
self.set_effect(None)
|
||||
self.set_view_options(None)
|
||||
self.temporal_transformer.reset_temp_vars()
|
||||
|
||||
def get_effect_mask(self, input_tensor: Tensor):
|
||||
@@ -459,7 +477,7 @@ class VanillaTemporalModule(nn.Module):
|
||||
|
||||
def forward(self, input_tensor: Tensor, encoder_hidden_states=None, attention_mask=None):
|
||||
if self.effect is None:
|
||||
return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)
|
||||
return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask, self.view_options)
|
||||
# return weighted average of input_tensor and AD output
|
||||
if type(self.effect) != Tensor:
|
||||
effect = self.effect
|
||||
@@ -468,7 +486,7 @@ class VanillaTemporalModule(nn.Module):
|
||||
return input_tensor
|
||||
else:
|
||||
effect = self.get_effect_mask(input_tensor)
|
||||
return input_tensor*(1.0-effect) + self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask)*effect
|
||||
return input_tensor*(1.0-effect) + self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask, self.view_options)*effect
|
||||
|
||||
|
||||
class TemporalTransformer3DModel(nn.Module):
|
||||
@@ -595,7 +613,7 @@ class TemporalTransformer3DModel(nn.Module):
|
||||
return self.temp_scale_mask[:, self.sub_idxs, :]
|
||||
return self.temp_scale_mask
|
||||
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
|
||||
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, view_options: ContextOptions=None):
|
||||
batch, channel, height, width = hidden_states.shape
|
||||
residual = hidden_states
|
||||
scale_mask = self.get_scale_mask(hidden_states)
|
||||
@@ -614,7 +632,8 @@ class TemporalTransformer3DModel(nn.Module):
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
video_length=self.video_length,
|
||||
scale_mask=scale_mask
|
||||
scale_mask=scale_mask,
|
||||
view_options=view_options
|
||||
)
|
||||
|
||||
# output
|
||||
@@ -691,26 +710,64 @@ class TemporalTransformerBlock(nn.Module):
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=None,
|
||||
video_length=None,
|
||||
scale_mask=None
|
||||
hidden_states: Tensor,
|
||||
encoder_hidden_states: Tensor=None,
|
||||
attention_mask: Tensor=None,
|
||||
video_length: int=None,
|
||||
scale_mask: Tensor=None,
|
||||
view_options: ContextOptions=None,
|
||||
):
|
||||
for attention_block, norm in zip(self.attention_blocks, self.norms):
|
||||
norm_hidden_states = norm(hidden_states).to(hidden_states.dtype)
|
||||
hidden_states = (
|
||||
attention_block(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states
|
||||
if attention_block.is_cross_attention
|
||||
else None,
|
||||
attention_mask=attention_mask,
|
||||
video_length=video_length,
|
||||
scale_mask=scale_mask
|
||||
if not view_options:
|
||||
for attention_block, norm in zip(self.attention_blocks, self.norms):
|
||||
norm_hidden_states = norm(hidden_states).to(hidden_states.dtype)
|
||||
hidden_states = (
|
||||
attention_block(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states
|
||||
if attention_block.is_cross_attention
|
||||
else None,
|
||||
attention_mask=attention_mask,
|
||||
video_length=video_length,
|
||||
scale_mask=scale_mask
|
||||
) + hidden_states
|
||||
)
|
||||
+ hidden_states
|
||||
)
|
||||
else:
|
||||
# views idea gotten from diffusers AnimateDiff FreeNoise implementation:
|
||||
# https://github.com/arthur-qiu/FreeNoise-AnimateDiff/blob/main/animatediff/models/motion_module.py
|
||||
# apply sliding context windows (views)
|
||||
views = get_context_windows(num_frames=video_length, opts=view_options)
|
||||
hidden_states = rearrange(hidden_states, "(b f) d c -> b f d c", f=video_length)
|
||||
value_final = torch.zeros_like(hidden_states)
|
||||
count_final = torch.zeros_like(hidden_states)
|
||||
batched_conds = hidden_states.size(1) // video_length
|
||||
for sub_idxs in views:
|
||||
weights = get_context_weights(len(sub_idxs), view_options.fuse_method) * batched_conds
|
||||
weights_tensor = torch.Tensor(weights).to(device=hidden_states.device).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
|
||||
|
||||
sub_hidden_states = rearrange(hidden_states[:, sub_idxs], "b f d c -> (b f) d c")
|
||||
for attention_block, norm in zip(self.attention_blocks, self.norms):
|
||||
norm_hidden_states = norm(sub_hidden_states).to(sub_hidden_states.dtype)
|
||||
sub_hidden_states = (
|
||||
attention_block(
|
||||
norm_hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states # do these need to be changed for sub_idxs too?
|
||||
if attention_block.is_cross_attention
|
||||
else None,
|
||||
attention_mask=attention_mask,
|
||||
video_length=len(sub_idxs),
|
||||
scale_mask=scale_mask[:, sub_idxs, :] if scale_mask is not None else scale_mask
|
||||
) + sub_hidden_states
|
||||
)
|
||||
sub_hidden_states = rearrange(sub_hidden_states, "(b f) d c -> b f d c", f=len(sub_idxs))
|
||||
|
||||
value_final[:, sub_idxs] += sub_hidden_states * weights_tensor
|
||||
count_final[:, sub_idxs] += weights_tensor
|
||||
|
||||
# get weighted average of sub_hidden_states
|
||||
hidden_states = value_final / count_final
|
||||
hidden_states = rearrange(hidden_states, "b f d c -> (b f) d c")
|
||||
del value_final
|
||||
del count_final
|
||||
|
||||
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
|
||||
|
||||
|
||||
+29
-18
@@ -5,7 +5,7 @@ import comfy.sample as comfy_sample
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .logger import logger
|
||||
from .model_utils import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path
|
||||
from .utils_model import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path
|
||||
from .motion_lora import MotionLoraInfo, MotionLoraList
|
||||
from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelSettings, load_motion_module
|
||||
from .sample_settings import SampleSettings, SeedNoiseGeneration
|
||||
@@ -15,7 +15,8 @@ from .nodes_gen1 import AnimateDiffLoaderWithContext
|
||||
from .nodes_gen2 import UseEvolvedSamplingNode, ApplyAnimateDiffModelNode, ApplyAnimateDiffModelBasicNode, LoadAnimateDiffModelNode, ADKeyframeNode
|
||||
from .nodes_multival import MultivalDynamicNode, MultivalFloatNode, MultivalScaledMaskNode
|
||||
from .nodes_sample import FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode
|
||||
from .nodes_context import LoopedUniformContextOptionsNode, StandardUniformContextOptionsNode, StandardStaticContextOptionsNode, BatchedContextOptionsNode
|
||||
from .nodes_context import (LoopedUniformContextOptionsNode, LoopedUniformViewOptionsNode, StandardUniformContextOptionsNode, StandardStaticContextOptionsNode, BatchedContextOptionsNode,
|
||||
StandardStaticViewOptionsNode, StandardUniformViewOptionsNode, ViewAsContextOptionsNode)
|
||||
from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect
|
||||
from .nodes_experimental import AnimateDiffModelSettingsSimple, AnimateDiffModelSettingsAdvanced, AnimateDiffModelSettingsAdvancedAttnStrengths
|
||||
from .nodes_deprecated import AnimateDiffLoader_Deprecated, AnimateDiffLoaderAdvanced_Deprecated, AnimateDiffCombine_Deprecated
|
||||
@@ -97,18 +98,23 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Multival Nodes
|
||||
"ADE_MultivalDynamic": MultivalDynamicNode,
|
||||
"ADE_MultivalScaledMask": MultivalScaledMaskNode,
|
||||
# Context Opts
|
||||
"ADE_StandardStaticContextOptions": StandardStaticContextOptionsNode,
|
||||
"ADE_StandardUniformContextOptions": StandardUniformContextOptionsNode,
|
||||
"ADE_AnimateDiffUniformContextOptions": LoopedUniformContextOptionsNode,
|
||||
"ADE_ViewsOnlyContextOptions": ViewAsContextOptionsNode,
|
||||
"ADE_BatchedContextOptions": BatchedContextOptionsNode,
|
||||
# View Opts
|
||||
"ADE_StandardStaticViewOptions": StandardStaticViewOptionsNode,
|
||||
"ADE_StandardUniformViewOptions": StandardUniformViewOptionsNode,
|
||||
"ADE_LoopedUniformViewOptions": LoopedUniformViewOptionsNode,
|
||||
# Iteration Opts
|
||||
"ADE_IterationOptsDefault": IterationOptionsNode,
|
||||
"ADE_IterationOptsFreeInit": FreeInitOptionsNode,
|
||||
# Noise Layer Nodes
|
||||
"ADE_NoiseLayerAdd": NoiseLayerAddNode,
|
||||
"ADE_NoiseLayerAddWeighted": NoiseLayerAddWeightedNode,
|
||||
"ADE_NoiseLayerReplace": NoiseLayerReplaceNode,
|
||||
# Context Opts
|
||||
"ADE_AnimateDiffUniformContextOptions": LoopedUniformContextOptionsNode,
|
||||
"ADE_StandardUniformContextOptions": StandardUniformContextOptionsNode,
|
||||
"ADE_StandardStaticContextOptions": StandardStaticContextOptionsNode,
|
||||
"ADE_BatchedContextOptions": BatchedContextOptionsNode,
|
||||
# Iteration Opts
|
||||
"ADE_IterationOptsDefault": IterationOptionsNode,
|
||||
"ADE_IterationOptsFreeInit": FreeInitOptionsNode,
|
||||
# Extras Nodes
|
||||
"ADE_AnimateDiffUnload": AnimateDiffUnload,
|
||||
"ADE_EmptyLatentImageLarge": EmptyLatentImageLarge,
|
||||
@@ -136,18 +142,23 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Multival Nodes
|
||||
"ADE_MultivalDynamic": "Multival Dynamic 🎭🅐🅓",
|
||||
"ADE_MultivalScaledMask": "Multival Scaled Mask 🎭🅐🅓",
|
||||
# Context Opts
|
||||
"ADE_StandardStaticContextOptions": "Context Options◆Standard Static 🎭🅐🅓",
|
||||
"ADE_StandardUniformContextOptions": "Context Options◆Standard Uniform 🎭🅐🅓",
|
||||
"ADE_AnimateDiffUniformContextOptions": "Context Options◆Looped Uniform 🎭🅐🅓",
|
||||
"ADE_ViewsOnlyContextOptions": "Context Options◆Views Only [VRAM⇈] 🎭🅐🅓",
|
||||
"ADE_BatchedContextOptions": "Context Options◆Batched [Non-AD] 🎭🅐🅓",
|
||||
# View Opts
|
||||
"ADE_StandardStaticViewOptions": "View Options◆Standard Static 🎭🅐🅓",
|
||||
"ADE_StandardUniformViewOptions": "View Options◆Standard Uniform 🎭🅐🅓",
|
||||
"ADE_LoopedUniformViewOptions": "View Options◆Looped Uniform 🎭🅐🅓",
|
||||
# Iteration Opts
|
||||
"ADE_IterationOptsDefault": "Default Iteration Options 🎭🅐🅓",
|
||||
"ADE_IterationOptsFreeInit": "FreeInit Iteration Options 🎭🅐🅓",
|
||||
# Noise Layer Nodes
|
||||
"ADE_NoiseLayerAdd": "Noise Layer [Add] 🎭🅐🅓",
|
||||
"ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] 🎭🅐🅓",
|
||||
"ADE_NoiseLayerReplace": "Noise Layer [Replace] 🎭🅐🅓",
|
||||
# Context Opts
|
||||
"ADE_AnimateDiffUniformContextOptions": "Looped Uniform Context Options 🎭🅐🅓",
|
||||
"ADE_StandardUniformContextOptions": "Standard Uniform Context Options 🎭🅐🅓",
|
||||
"ADE_StandardStaticContextOptions": "Standard Static Context Options 🎭🅐🅓",
|
||||
"ADE_BatchedContextOptions": "[Non-AD] Batched Context Options 🎭🅐🅓",
|
||||
# Iteration Opts
|
||||
"ADE_IterationOptsDefault": "Default Iteration Options 🎭🅐🅓",
|
||||
"ADE_IterationOptsFreeInit": "FreeInit Iteration Options 🎭🅐🅓",
|
||||
# Extras Nodes
|
||||
"ADE_AnimateDiffUnload": "AnimateDiff Unload 🎭🅐🅓",
|
||||
"ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch) 🎭🅐🅓",
|
||||
|
||||
+217
-18
@@ -1,4 +1,10 @@
|
||||
from .context import ContextFuseMethod, ContextOptions, ContextSchedules
|
||||
from .context import ContextFuseMethod, ContextOptions, ContextOptionsGroup, ContextSchedules
|
||||
from .utils_model import BIGMAX
|
||||
|
||||
|
||||
LENGTH_MAX = 128 # keep an eye on these max values;
|
||||
STRIDE_MAX = 32 # would need to be updated
|
||||
OVERLAP_MAX = 128 # if new motion modules come out
|
||||
|
||||
|
||||
class LoopedUniformContextOptionsNode:
|
||||
@@ -6,24 +12,35 @@ class LoopedUniformContextOptionsNode:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"context_length": ("INT", {"default": 16, "min": 1, "max": 128}), # keep an eye on these max values
|
||||
"context_stride": ("INT", {"default": 1, "min": 1, "max": 32}), # would need to be updated
|
||||
"context_overlap": ("INT", {"default": 4, "min": 0, "max": 128}), # if new motion modules come out
|
||||
"context_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}),
|
||||
"context_stride": ("INT", {"default": 1, "min": 1, "max": STRIDE_MAX}),
|
||||
"context_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}),
|
||||
"context_schedule": (ContextSchedules.UNIFORM_SCHEDULE_LIST,),
|
||||
"closed_loop": ("BOOLEAN", {"default": False},),
|
||||
#"sync_context_to_pe": ("BOOLEAN", {"default": False},),
|
||||
},
|
||||
"optional": {
|
||||
"fuse_method": (ContextFuseMethod.LIST,),
|
||||
"use_on_equal_length": ("BOOLEAN", {"default": False},),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
"prev_context": ("CONTEXT_OPTIONS",),
|
||||
"view_opts": ("VIEW_OPTS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXT_OPTIONS",)
|
||||
RETURN_NAMES = ("CONTEXT_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts"
|
||||
FUNCTION = "create_options"
|
||||
|
||||
def create_options(self, context_length: int, context_stride: int, context_overlap: int, context_schedule: int, closed_loop: bool,
|
||||
fuse_method: str=ContextFuseMethod.FLAT):
|
||||
fuse_method: str=ContextFuseMethod.FLAT, use_on_equal_length=False, start_percent: float=0.0, guarantee_steps: int=1,
|
||||
view_opts: ContextOptions=None, prev_context: ContextOptionsGroup=None):
|
||||
if prev_context is None:
|
||||
prev_context = ContextOptionsGroup()
|
||||
prev_context = prev_context.clone()
|
||||
|
||||
context_options = ContextOptions(
|
||||
context_length=context_length,
|
||||
context_stride=context_stride,
|
||||
@@ -31,9 +48,14 @@ class LoopedUniformContextOptionsNode:
|
||||
context_schedule=context_schedule,
|
||||
closed_loop=closed_loop,
|
||||
fuse_method=fuse_method,
|
||||
use_on_equal_length=use_on_equal_length,
|
||||
start_percent=start_percent,
|
||||
guarantee_steps=guarantee_steps,
|
||||
view_options=view_opts,
|
||||
)
|
||||
#context_options.set_sync_context_to_pe(sync_context_to_pe)
|
||||
return (context_options,)
|
||||
prev_context.add(context_options)
|
||||
return (prev_context,)
|
||||
|
||||
|
||||
class StandardUniformContextOptionsNode:
|
||||
@@ -41,21 +63,32 @@ class StandardUniformContextOptionsNode:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"context_length": ("INT", {"default": 16, "min": 1, "max": 128}), # keep an eye on these max values
|
||||
"context_stride": ("INT", {"default": 1, "min": 1, "max": 32}), # would need to be updated
|
||||
"context_overlap": ("INT", {"default": 4, "min": 0, "max": 128}), # if new motion modules come out
|
||||
"context_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}),
|
||||
"context_stride": ("INT", {"default": 1, "min": 1, "max": STRIDE_MAX}),
|
||||
"context_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}),
|
||||
},
|
||||
"optional": {
|
||||
"fuse_method": (ContextFuseMethod.LIST,),
|
||||
"use_on_equal_length": ("BOOLEAN", {"default": False},),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
"prev_context": ("CONTEXT_OPTIONS",),
|
||||
"view_opts": ("VIEW_OPTS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXT_OPTIONS",)
|
||||
RETURN_NAMES = ("CONTEXT_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts"
|
||||
FUNCTION = "create_options"
|
||||
|
||||
def create_options(self, context_length: int, context_stride: int, context_overlap: int,
|
||||
fuse_method: str=ContextFuseMethod.FLAT):
|
||||
fuse_method: str=ContextFuseMethod.FLAT, use_on_equal_length=False, start_percent: float=0.0, guarantee_steps: int=1,
|
||||
view_opts: ContextOptions=None, prev_context: ContextOptionsGroup=None):
|
||||
if prev_context is None:
|
||||
prev_context = ContextOptionsGroup()
|
||||
prev_context = prev_context.clone()
|
||||
|
||||
context_options = ContextOptions(
|
||||
context_length=context_length,
|
||||
context_stride=context_stride,
|
||||
@@ -63,8 +96,13 @@ class StandardUniformContextOptionsNode:
|
||||
context_schedule=ContextSchedules.UNIFORM_STANDARD,
|
||||
closed_loop=False,
|
||||
fuse_method=fuse_method,
|
||||
use_on_equal_length=use_on_equal_length,
|
||||
start_percent=start_percent,
|
||||
guarantee_steps=guarantee_steps,
|
||||
view_options=view_opts,
|
||||
)
|
||||
return (context_options,)
|
||||
prev_context.add(context_options)
|
||||
return (prev_context,)
|
||||
|
||||
|
||||
class StandardStaticContextOptionsNode:
|
||||
@@ -72,28 +110,44 @@ class StandardStaticContextOptionsNode:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"context_length": ("INT", {"default": 16, "min": 1, "max": 128}), # keep an eye on these max values
|
||||
"context_overlap": ("INT", {"default": 4, "min": 0, "max": 128}), # if new motion modules come out
|
||||
"context_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}),
|
||||
"context_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}),
|
||||
},
|
||||
"optional": {
|
||||
"fuse_method": (ContextFuseMethod.LIST, {"default": ContextFuseMethod.PYRAMID}),
|
||||
"use_on_equal_length": ("BOOLEAN", {"default": False},),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
"prev_context": ("CONTEXT_OPTIONS",),
|
||||
"view_opts": ("VIEW_OPTS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXT_OPTIONS",)
|
||||
RETURN_NAMES = ("CONTEXT_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts"
|
||||
FUNCTION = "create_options"
|
||||
|
||||
def create_options(self, context_length: int, context_overlap: int,
|
||||
fuse_method: str=ContextFuseMethod.FLAT):
|
||||
fuse_method: str=ContextFuseMethod.FLAT, use_on_equal_length=False, start_percent: float=0.0, guarantee_steps: int=1,
|
||||
view_opts: ContextOptions=None, prev_context: ContextOptionsGroup=None):
|
||||
if prev_context is None:
|
||||
prev_context = ContextOptionsGroup()
|
||||
prev_context = prev_context.clone()
|
||||
|
||||
context_options = ContextOptions(
|
||||
context_length=context_length,
|
||||
context_stride=None,
|
||||
context_overlap=context_overlap,
|
||||
context_schedule=ContextSchedules.STATIC_STANDARD,
|
||||
fuse_method=fuse_method,
|
||||
use_on_equal_length=use_on_equal_length,
|
||||
start_percent=start_percent,
|
||||
guarantee_steps=guarantee_steps,
|
||||
view_options=view_opts,
|
||||
)
|
||||
return (context_options,)
|
||||
prev_context.add(context_options)
|
||||
return (prev_context,)
|
||||
|
||||
|
||||
class BatchedContextOptionsNode:
|
||||
@@ -101,18 +155,163 @@ class BatchedContextOptionsNode:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"context_length": ("INT", {"default": 16, "min": 1, "max": 128}), # keep an eye on these max values
|
||||
"context_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}),
|
||||
},
|
||||
"optional": {
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
"prev_context": ("CONTEXT_OPTIONS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXT_OPTIONS",)
|
||||
RETURN_NAMES = ("CONTEXT_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts"
|
||||
FUNCTION = "create_options"
|
||||
|
||||
def create_options(self, context_length: int):
|
||||
def create_options(self, context_length: int, start_percent: float=0.0, guarantee_steps: int=1,
|
||||
prev_context: ContextOptionsGroup=None):
|
||||
if prev_context is None:
|
||||
prev_context = ContextOptionsGroup()
|
||||
prev_context = prev_context.clone()
|
||||
|
||||
context_options = ContextOptions(
|
||||
context_length=context_length,
|
||||
context_overlap=0,
|
||||
context_schedule=ContextSchedules.BATCHED,
|
||||
start_percent=start_percent,
|
||||
guarantee_steps=guarantee_steps,
|
||||
)
|
||||
return (context_options,)
|
||||
prev_context.add(context_options)
|
||||
return (prev_context,)
|
||||
|
||||
|
||||
class ViewAsContextOptionsNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"view_opts_req": ("VIEW_OPTS",),
|
||||
},
|
||||
"optional": {
|
||||
"use_on_equal_length": ("BOOLEAN", {"default": False},),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
"prev_context": ("CONTEXT_OPTIONS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONTEXT_OPTIONS",)
|
||||
RETURN_NAMES = ("CONTEXT_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts"
|
||||
FUNCTION = "create_options"
|
||||
|
||||
def create_options(self, view_opts_req: ContextOptions, start_percent: float=0.0, guarantee_steps: int=1,
|
||||
prev_context: ContextOptionsGroup=None):
|
||||
if prev_context is None:
|
||||
prev_context = ContextOptionsGroup()
|
||||
prev_context = prev_context.clone()
|
||||
context_options = ContextOptions(
|
||||
context_schedule=ContextSchedules.VIEW_AS_CONTEXT,
|
||||
start_percent=start_percent,
|
||||
guarantee_steps=guarantee_steps,
|
||||
view_options=view_opts_req,
|
||||
use_on_equal_length=True
|
||||
)
|
||||
prev_context.add(context_options)
|
||||
return (prev_context,)
|
||||
|
||||
|
||||
#########################
|
||||
# View Options
|
||||
class StandardStaticViewOptionsNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"view_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}),
|
||||
"view_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}),
|
||||
},
|
||||
"optional": {
|
||||
"fuse_method": (ContextFuseMethod.LIST, {"default": ContextFuseMethod.PYRAMID}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VIEW_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/view opts"
|
||||
FUNCTION = "create_options"
|
||||
|
||||
def create_options(self, view_length: int, view_overlap: int,
|
||||
fuse_method: str=ContextFuseMethod.FLAT,):
|
||||
view_options = ContextOptions(
|
||||
context_length=view_length,
|
||||
context_stride=None,
|
||||
context_overlap=view_overlap,
|
||||
context_schedule=ContextSchedules.STATIC_STANDARD,
|
||||
fuse_method=fuse_method,
|
||||
)
|
||||
return (view_options,)
|
||||
|
||||
|
||||
class StandardUniformViewOptionsNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"view_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}),
|
||||
"view_stride": ("INT", {"default": 1, "min": 1, "max": STRIDE_MAX}),
|
||||
"view_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}),
|
||||
},
|
||||
"optional": {
|
||||
"fuse_method": (ContextFuseMethod.LIST, {"default": ContextFuseMethod.PYRAMID}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VIEW_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/view opts"
|
||||
FUNCTION = "create_options"
|
||||
|
||||
def create_options(self, view_length: int, view_overlap: int, view_stride: int,
|
||||
fuse_method: str=ContextFuseMethod.FLAT,):
|
||||
view_options = ContextOptions(
|
||||
context_length=view_length,
|
||||
context_stride=view_stride,
|
||||
context_overlap=view_overlap,
|
||||
context_schedule=ContextSchedules.UNIFORM_STANDARD,
|
||||
fuse_method=fuse_method,
|
||||
)
|
||||
return (view_options,)
|
||||
|
||||
|
||||
class LoopedUniformViewOptionsNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"view_length": ("INT", {"default": 16, "min": 1, "max": LENGTH_MAX}),
|
||||
"view_stride": ("INT", {"default": 1, "min": 1, "max": STRIDE_MAX}),
|
||||
"view_overlap": ("INT", {"default": 4, "min": 0, "max": OVERLAP_MAX}),
|
||||
"closed_loop": ("BOOLEAN", {"default": False},),
|
||||
},
|
||||
"optional": {
|
||||
"fuse_method": (ContextFuseMethod.LIST, {"default": ContextFuseMethod.PYRAMID}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("VIEW_OPTS",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/context opts/view opts"
|
||||
FUNCTION = "create_options"
|
||||
|
||||
def create_options(self, view_length: int, view_overlap: int, view_stride: int, closed_loop: bool,
|
||||
fuse_method: str=ContextFuseMethod.FLAT,):
|
||||
view_options = ContextOptions(
|
||||
context_length=view_length,
|
||||
context_stride=view_stride,
|
||||
context_overlap=view_overlap,
|
||||
context_schedule=ContextSchedules.UNIFORM_LOOPED,
|
||||
closed_loop=closed_loop,
|
||||
fuse_method=fuse_method,
|
||||
)
|
||||
return (view_options,)
|
||||
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .context import ContextSchedules, ContextOptions
|
||||
from .logger import logger
|
||||
from .model_utils import Folders, BetaSchedules, get_available_motion_models
|
||||
from .utils_model import Folders, BetaSchedules, get_available_motion_models
|
||||
from .model_injection import ModelPatcherAndInjector, InjectionParams, MotionModelGroup, load_motion_module
|
||||
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from comfy.model_patcher import ModelPatcher
|
||||
from comfy.sd import load_checkpoint_guess_config
|
||||
|
||||
from .logger import logger
|
||||
from .model_utils import IsChangedHelper, BetaSchedules
|
||||
from .utils_model import IsChangedHelper, BetaSchedules
|
||||
from .model_injection import get_vanilla_model_patcher
|
||||
|
||||
|
||||
|
||||
@@ -4,10 +4,10 @@ import torch
|
||||
import comfy.sample as comfy_sample
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .context import ContextOptions, ContextSchedules
|
||||
from .context import ContextOptions, ContextOptionsGroup, ContextSchedules
|
||||
from .logger import logger
|
||||
from .model_utils import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path
|
||||
from .motion_utils import ADKeyframeGroup
|
||||
from .utils_model import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path
|
||||
from .utils_motion import ADKeyframeGroup
|
||||
from .motion_lora import MotionLoraInfo, MotionLoraList
|
||||
from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelSettings, load_motion_module
|
||||
from .sample_settings import SampleSettings, SeedNoiseGeneration
|
||||
@@ -43,7 +43,7 @@ class AnimateDiffLoaderWithContext:
|
||||
def load_mm_and_inject_params(self,
|
||||
model: ModelPatcher,
|
||||
model_name: str, beta_schedule: str,# apply_mm_groupnorm_hack: bool,
|
||||
context_options: ContextOptions=None, motion_lora: MotionLoraList=None, motion_model_settings: MotionModelSettings=None,
|
||||
context_options: ContextOptionsGroup=None, motion_lora: MotionLoraList=None, motion_model_settings: MotionModelSettings=None,
|
||||
sample_settings: SampleSettings=None, motion_scale: float=1.0, apply_v2_models_properly: bool=False, ad_keyframes: ADKeyframeGroup=None,
|
||||
):
|
||||
# load motion module
|
||||
|
||||
@@ -4,10 +4,10 @@ import torch
|
||||
import comfy.sample as comfy_sample
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from .context import ContextOptions, ContextSchedules
|
||||
from .context import ContextOptions, ContextOptionsGroup, ContextSchedules
|
||||
from .logger import logger
|
||||
from .model_utils import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path
|
||||
from .motion_utils import ADKeyframeGroup, ADKeyframe
|
||||
from .utils_model import BetaSchedules, get_available_motion_loras, get_available_motion_models, get_motion_lora_path
|
||||
from .utils_motion import ADKeyframeGroup, ADKeyframe
|
||||
from .motion_lora import MotionLoraInfo, MotionLoraList
|
||||
from .model_injection import (InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, MotionModelSettings,
|
||||
load_motion_module, load_motion_module_gen2, load_motion_lora_as_patches, validate_model_compatibility_gen2)
|
||||
@@ -35,7 +35,7 @@ class UseEvolvedSamplingNode:
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓"
|
||||
FUNCTION = "use_evolved_sampling"
|
||||
|
||||
def use_evolved_sampling(self, model: ModelPatcher, beta_schedule: str, m_models: MotionModelGroup=None, context_options: ContextOptions=None,
|
||||
def use_evolved_sampling(self, model: ModelPatcher, beta_schedule: str, m_models: MotionModelGroup=None, context_options: ContextOptionsGroup=None,
|
||||
sample_settings: SampleSettings=None, beta_schedule_override=None):
|
||||
if m_models is not None:
|
||||
m_models = m_models.clone()
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Union
|
||||
import torch
|
||||
from torch import Tensor
|
||||
|
||||
from .motion_utils import linear_conversion, normalize_min_max
|
||||
from .utils_motion import linear_conversion, normalize_min_max
|
||||
|
||||
|
||||
class ScaleType:
|
||||
|
||||
@@ -2,7 +2,7 @@ from torch import Tensor
|
||||
|
||||
from .freeinit import FreeInitFilter
|
||||
from .sample_settings import FreeInitOptions, IterationOptions, NoiseLayerAdd, NoiseLayerAddWeighted, NoiseLayerGroup, NoiseLayerReplace, NoiseLayerType, SeedNoiseGeneration, SampleSettings
|
||||
from .model_utils import BIGMIN, BIGMAX
|
||||
from .utils_model import BIGMIN, BIGMAX
|
||||
|
||||
|
||||
class SampleSettingsNode:
|
||||
|
||||
@@ -7,7 +7,7 @@ import comfy.samplers
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
|
||||
from . import freeinit
|
||||
from .context import ContextOptions
|
||||
from .context import ContextOptions, ContextOptionsGroup
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -273,8 +273,8 @@ class SeedNoiseGeneration:
|
||||
@staticmethod
|
||||
def _convert_to_repeated_context(noise: Tensor, extra_args: dict, **kwargs):
|
||||
# if no context_length, return unmodified noise
|
||||
opts: ContextOptions = extra_args["context_options"]
|
||||
context_length: int = opts.context_length
|
||||
opts: ContextOptionsGroup = extra_args["context_options"]
|
||||
context_length: int = opts.context_length if not opts.view_options else opts.view_options.context_length
|
||||
if context_length is None:
|
||||
return noise
|
||||
length = noise.shape[0]
|
||||
@@ -285,9 +285,9 @@ class SeedNoiseGeneration:
|
||||
@staticmethod
|
||||
def _convert_to_freenoise(noise: Tensor, seed: int, extra_args: dict, **kwargs):
|
||||
# if no context_length, return unmodified noise
|
||||
opts: ContextOptions = extra_args["context_options"]
|
||||
context_length: int = opts.context_length
|
||||
context_overlap: int = opts.context_overlap
|
||||
opts: ContextOptionsGroup = extra_args["context_options"]
|
||||
context_length: int = opts.context_length if not opts.view_options else opts.view_options.context_length
|
||||
context_overlap: int = opts.context_overlap if not opts.view_options else opts.view_options.context_overlap
|
||||
video_length: int = noise.shape[0]
|
||||
if context_length is None:
|
||||
return noise
|
||||
|
||||
+32
-23
@@ -14,10 +14,10 @@ import comfy.sample
|
||||
import comfy.utils
|
||||
from comfy.controlnet import ControlBase
|
||||
|
||||
from .context import ContextFuseMethod, generate_distance_weight, get_context_windows
|
||||
from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows
|
||||
from .sample_settings import IterationOptions, SeedNoiseGeneration, prepare_mask_ad
|
||||
from .motion_utils import GroupNormAD
|
||||
from .model_utils import ModelTypeSD, wrap_function_to_inject_xformers_bug_info
|
||||
from .utils_motion import GroupNormAD
|
||||
from .utils_model import ModelTypeSD, wrap_function_to_inject_xformers_bug_info
|
||||
from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher
|
||||
from .motion_module_ad import AnimateDiffFormat, AnimateDiffInfo, AnimateDiffVersion, VanillaTemporalModule
|
||||
from .logger import logger
|
||||
@@ -115,13 +115,13 @@ def groupnorm_mm_factory(params: InjectionParams):
|
||||
# axes_factor normalizes batch based on total conds and unconds passed in batch;
|
||||
# the conds and unconds per batch can change based on VRAM optimizations that may kick in
|
||||
if not params.is_using_sliding_context():
|
||||
axes_factor = input.size(0)//params.full_length
|
||||
batched_conds = input.size(0)//params.full_length
|
||||
else:
|
||||
axes_factor = input.size(0)//params.context_options.context_length
|
||||
batched_conds = input.size(0)//params.context_options.context_length
|
||||
|
||||
input = rearrange(input, "(b f) c h w -> b c f h w", b=axes_factor)
|
||||
input = rearrange(input, "(b f) c h w -> b c f h w", b=batched_conds)
|
||||
input = group_norm(input, self.num_groups, self.weight, self.bias, self.eps)
|
||||
input = rearrange(input, "b c f h w -> (b f) c h w", b=axes_factor)
|
||||
input = rearrange(input, "b c f h w -> (b f) c h w", b=batched_conds)
|
||||
return input
|
||||
return groupnorm_mm_forward
|
||||
|
||||
@@ -141,7 +141,15 @@ def get_additional_models_factory(orig_get_additional_models: Callable, motion_m
|
||||
|
||||
def apply_params_to_motion_models(motion_models: MotionModelGroup, params: InjectionParams):
|
||||
params = params.clone()
|
||||
if params.context_options.context_length and params.full_length > params.context_options.context_length:
|
||||
if params.context_options.context_schedule == ContextSchedules.VIEW_AS_CONTEXT:
|
||||
params.context_options._current_context.context_length = params.full_length
|
||||
# check (and message) should be different based on use_on_equal_length setting
|
||||
if params.context_options.context_length:
|
||||
pass
|
||||
|
||||
allow_equal = params.context_options.use_on_equal_length
|
||||
enough_latents = params.full_length >= params.context_options.context_length if allow_equal else params.full_length > params.context_options.context_length
|
||||
if params.context_options.context_length and enough_latents:
|
||||
logger.info(f"Sliding context window activated - latents passed in ({params.full_length}) greater than context_length {params.context_options.context_length}.")
|
||||
else:
|
||||
logger.info(f"Regular AnimateDiff activated - latents passed in ({params.full_length}) less or equal to context_length {params.context_options.context_length}.")
|
||||
@@ -156,7 +164,9 @@ def apply_params_to_motion_models(motion_models: MotionModelGroup, params: Injec
|
||||
# otherwise, treat context_length as intended AD frame window
|
||||
else:
|
||||
for motion_model in motion_models.models:
|
||||
if params.context_options.context_length > motion_model.model.encoding_max_len:
|
||||
view_options = params.context_options.view_options
|
||||
context_length = view_options.context_length if view_options else params.context_options.context_length
|
||||
if context_length > motion_model.model.encoding_max_len:
|
||||
raise ValueError(f"AnimateDiff model {motion_model.model.mm_info.mm_name} has upper limit of {motion_model.model.encoding_max_len} frames for a context window, but received context length of {params.context_options.context_length}.")
|
||||
motion_models.set_video_length(params.context_options.context_length, params.full_length)
|
||||
# inject model
|
||||
@@ -422,9 +432,13 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep,
|
||||
return resized_cond
|
||||
|
||||
# get context windows
|
||||
context_windows = get_context_windows(ADGS.current_step, ADGS.params.full_length, ADGS.params.context_options)
|
||||
ADGS.params.context_options.step = ADGS.current_step
|
||||
context_windows = get_context_windows(ADGS.params.full_length, ADGS.params.context_options)
|
||||
# figure out how input is split
|
||||
axes_factor = x_in.size(0)//ADGS.params.full_length
|
||||
batched_conds = x_in.size(0)//ADGS.params.full_length
|
||||
|
||||
if ADGS.motion_models is not None:
|
||||
ADGS.motion_models.set_view_options(ADGS.params.context_options.view_options)
|
||||
|
||||
# prepare final cond, uncond, and out_count
|
||||
cond_final = torch.zeros_like(x_in)
|
||||
@@ -442,7 +456,7 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep,
|
||||
model_options["transformer_options"]["ad_params"]["context_length"] = len(ctx_idxs)
|
||||
# account for all portions of input frames
|
||||
full_idxs = []
|
||||
for n in range(axes_factor):
|
||||
for n in range(batched_conds):
|
||||
for ind in ctx_idxs:
|
||||
full_idxs.append((ADGS.params.full_length*n)+ind)
|
||||
# get subsections of x, timestep, cond, uncond, cond_concat
|
||||
@@ -453,17 +467,12 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep,
|
||||
|
||||
sub_cond_out, sub_uncond_out = comfy.samplers.calc_cond_uncond_batch(model, sub_cond, sub_uncond, sub_x, sub_timestep, model_options)
|
||||
|
||||
if ADGS.params.context_options.fuse_method == ContextFuseMethod.FLAT:
|
||||
# equal weights for idxs
|
||||
cond_final[full_idxs] += sub_cond_out
|
||||
uncond_final[full_idxs] += sub_uncond_out
|
||||
out_count_final[full_idxs] += 1 # increment which indeces were used
|
||||
elif ADGS.params.context_options.fuse_method == ContextFuseMethod.PYRAMID:
|
||||
# greater weight towards center of idxs
|
||||
weights = torch.Tensor(generate_distance_weight(len(ctx_idxs)) * axes_factor).to(device=x_in.device).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
|
||||
cond_final[full_idxs] += sub_cond_out * weights
|
||||
uncond_final[full_idxs] += sub_uncond_out * weights
|
||||
out_count_final[full_idxs] += weights
|
||||
# add conds and counts based on weights of fuse method
|
||||
weights = get_context_weights(len(ctx_idxs), ADGS.params.context_options.fuse_method) * batched_conds
|
||||
weights_tensor = torch.Tensor(weights).to(device=x_in.device).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
|
||||
cond_final[full_idxs] += sub_cond_out * weights_tensor
|
||||
uncond_final[full_idxs] += sub_uncond_out * weights_tensor
|
||||
out_count_final[full_idxs] += weights_tensor
|
||||
|
||||
|
||||
# normalize cond and uncond via division by context usage counts
|
||||
|
||||
@@ -27,7 +27,6 @@ else:
|
||||
optimized_attention_mm = attention_sub_quad
|
||||
|
||||
|
||||
# maintain backwards compatibility with the comfy.ops hasattr check (TODO: remove once a non-backwards compatible change happens)
|
||||
class CrossAttentionMM(nn.Module):
|
||||
def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0., dtype=None, device=None,
|
||||
operations=comfy.ops.disable_weight_init):
|
||||
@@ -110,6 +109,27 @@ def extend_to_batch_size(tensor: Tensor, batch_size: int):
|
||||
return tensor
|
||||
|
||||
|
||||
def get_sorted_list_via_attr(objects: list, attr: str) -> list:
|
||||
if not objects:
|
||||
return objects
|
||||
elif len(objects) <= 1:
|
||||
return [x for x in objects]
|
||||
# now that we know we have to sort, do it following these rules:
|
||||
# a) if objects have same value of attribute, maintain their relative order
|
||||
# b) perform sorting of the groups of objects with same attributes
|
||||
unique_attrs = {}
|
||||
for object in objects:
|
||||
val_attr = getattr(objects, attr)
|
||||
unique_attrs.get(val_attr, list()).append(object)
|
||||
# now that we have the unique attr values grouped together in relative order, sort them by key
|
||||
sorted_attrs = dict(sorted(unique_attrs.items()))
|
||||
# now flatten out the dict into a list to return
|
||||
sorted_list = []
|
||||
for object_list in sorted_attrs.values():
|
||||
sorted_list.extend(object_list)
|
||||
return sorted_list
|
||||
|
||||
|
||||
class MotionCompatibilityError(ValueError):
|
||||
pass
|
||||
|
||||
Reference in New Issue
Block a user