Views, start of Context scheduling, use_on_equal_length for Contexts,

This commit is contained in:
Jedrzej Kosinski
2024-01-13 10:28:59 -06:00
parent 666b5f047d
commit dcfbee1454
16 changed files with 567 additions and 294 deletions
+1 -1
View File
@@ -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
View File
@@ -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 -8
View File
@@ -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(
+81 -24
View File
@@ -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
View File
@@ -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
View File
@@ -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,)
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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 -4
View File
@@ -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 -4
View File
@@ -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()
+1 -1
View File
@@ -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:
+1 -1
View File
@@ -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:
+6 -6
View File
@@ -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
View File
@@ -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