Added AnimateDiff-SDXL support - you need to use 'linear (AnimateDiff-SDXL)' beta_schedule

This commit is contained in:
Jedrzej Kosinski
2023-11-10 01:24:59 -06:00
parent 859c1fe2a7
commit 2d2a5c39bb
7 changed files with 265 additions and 13 deletions
+8 -2
View File
@@ -29,15 +29,17 @@ class ModelSamplingConfig:
class BetaSchedules:
SQRT_LINEAR = "sqrt_linear (AnimateDiff)"
LINEAR_ADXL = "linear (AnimateDiff-SDXL)"
LINEAR = "linear (HotshotXL/default)"
SQRT = "sqrt"
COSINE = "cosine"
SQUAREDCOS_CAP_V2 = "squaredcos_cap_v2"
ALIAS_LIST = [SQRT_LINEAR, LINEAR, SQRT, COSINE, SQUAREDCOS_CAP_V2]
ALIAS_LIST = [SQRT_LINEAR, LINEAR_ADXL, LINEAR, SQRT, COSINE, SQUAREDCOS_CAP_V2]
ALIAS_MAP = {
SQRT_LINEAR: "sqrt_linear",
LINEAR_ADXL: "linear", # also linear, but has different linear_end (0.020)
LINEAR: "linear",
SQRT: "sqrt",
COSINE: "cosine",
@@ -54,7 +56,11 @@ class BetaSchedules:
@classmethod
def to_model_sampling(cls, alias: str, model: ModelPatcher):
return model_sampling(cls.to_config(alias), model_type=model.model.model_type)
ms_obj = model_sampling(cls.to_config(alias), model_type=model.model.model_type)
if alias == cls.LINEAR_ADXL:
# uses linear_end=0.020
ms_obj._register_schedule(given_betas=None, beta_schedule=cls.to_name(alias), timesteps=1000, linear_start=0.00085, linear_end=0.020, cosine_s=8e-3)
return ms_obj
@staticmethod
def get_alias_list_with_first_element(first_element: str):
+93 -5
View File
@@ -13,9 +13,10 @@ from .logger import logger
from .model_utils import ModelTypesSD, calculate_file_hash, get_motion_lora_path, get_motion_model_path, \
get_sd_model_type
from .motion_lora import MotionLoRAList, MotionLoRAWrapper
from .motion_module_ad import AnimDiffMotionWrapper, has_mid_block
from .motion_module_ad import AnimDiffMotionWrapper, VanillaTemporalModule, has_mid_block
from .motion_module_adxl import AnimDiffSDXLMotionWrapper
from .motion_module_hsxl import HotShotXLMotionWrapper, TransformerTemporal
from .motion_utils import GenericMotionWrapper, InjectorVersion, NoiseType, normalize_min_max
from .motion_utils import GenericMotionWrapper, InjectorVersion, MotionCompatibilityError, NoiseType, normalize_min_max
# inject into ModelPatcher.clone to carry over injected params over to cloned ModelPatcher
orig_modelpatcher_clone = comfy_model_patcher.ModelPatcher.clone
@@ -225,13 +226,19 @@ def load_motion_module(model_name: str, motion_lora: MotionLoRAList = None, mode
if sd_model_type == ModelTypesSD.SD1_5:
try:
motion_module = AnimDiffMotionWrapper(mm_state_dict=mm_state_dict, mm_hash=model_hash, mm_name=model_name, loras=loras)
except ValueError as e:
except MotionCompatibilityError as e:
raise ValueError(f"Motion model {model_name} is not compatible with SD1.5-based model.", e)
elif sd_model_type == ModelTypesSD.SDXL:
# determine if motion module is a AnimateDiffXL model or a HotshotXL model
try:
# try to load as HotShotXL model first
motion_module = HotShotXLMotionWrapper(mm_state_dict=mm_state_dict, mm_hash=model_hash, mm_name=model_name, loras=loras)
except ValueError as e:
raise ValueError(f"Motion model {model_name} is not compatible with SDXL-based model.", e)
except MotionCompatibilityError as e:
# if not compatible, try to load as AnimateDiff-SDXL
try:
motion_module = AnimDiffSDXLMotionWrapper(mm_state_dict=mm_state_dict, mm_hash=model_hash, mm_name=model_name, loras=loras)
except MotionCompatibilityError as e:
raise ValueError(f"Motion model {model_name} is not compatible with neither AnimateDiff-SDXL nor HotShotXL.", e)
else:
raise ValueError(f"SD model must be either SD1.5-based for AnimateDiff or SDXL-based for HotShotXL.")
@@ -361,6 +368,85 @@ def _eject_motion_module_from_unet(model: ModelPatcher):
############################################################################################################
############################################################################################################
## AnimateDiff-SDXL
def _inject_adxl_motion_module_to_unet(model: ModelPatcher, motion_module: 'AnimDiffSDXLMotionWrapper'):
unet: openaimodel.UNetModel = model.model.diffusion_model
# inject input (down) blocks
# AnimateDiffSDXL mm contains 3 downblocks, each with 2 TransformerTemporals - 6 in total
# per_block is the amount of Temporal Blocks per down block
_perform_adxl_motion_module_injection(unet.input_blocks, motion_module.down_blocks, injection_goal=6, per_block=2)
# inject output (up) blocks
# AnimateDiffSDXL mm contains 3 upblocks, each with 3 TransformerTemporals - 9 in total
_perform_adxl_motion_module_injection(unet.output_blocks, motion_module.up_blocks, injection_goal=9, per_block=3)
# inject mid block, if needed (encapsulate in list to make structure compatible)
if motion_module.mid_block is not None:
_perform_adxl_motion_module_injection(unet.middle_block, [motion_module.mid_block], injection_goal=1, per_block=1)
# keep track of if unet blocks actually affected
set_injected_unet_version(model, InjectorVersion.ADXL_V1_V2)
def _perform_adxl_motion_module_injection(unet_blocks: nn.ModuleList, mm_blocks: nn.ModuleList, injection_goal: int, per_block: int):
# Rules for injection:
# For each component list in a unet block:
# if SpatialTransformer exists in list, place next block after last occurrence
# elif ResBlock exists in list, place next block after first occurrence
# else don't place block
injection_count = 0
unet_idx = 0
# only stop injecting when modules exhausted
while injection_count < injection_goal:
# figure out which VanillaTemporalModule from mm to inject
mm_blk_idx, mm_vtm_idx = injection_count // per_block, injection_count % per_block
# figure out layout of unet block components
st_idx = -1 # SpatialTransformer index
res_idx = -1 # first ResBlock index
# first, figure out indeces of relevant blocks
for idx, component in enumerate(unet_blocks[unet_idx]):
if type(component) == SpatialTransformer:
st_idx = idx
elif type(component) == ResBlock and res_idx < 0:
res_idx = idx
# if SpatialTransformer exists, inject right after
if st_idx >= 0:
#logger.info(f"ADXL: injecting after ST({st_idx})")
unet_blocks[unet_idx].insert(st_idx+1, mm_blocks[mm_blk_idx].motion_modules[mm_vtm_idx])
injection_count += 1
# otherwise, if only ResBlock exists, inject right after
elif res_idx >= 0:
#logger.info(f"ADXL: injecting after Res({res_idx})")
unet_blocks[unet_idx].insert(res_idx+1, mm_blocks[mm_blk_idx].motion_modules[mm_vtm_idx])
injection_count += 1
# increment unet_idx
unet_idx += 1
def _eject_adxl_motion_module_from_unet(model: ModelPatcher):
unet: openaimodel.UNetModel = model.model.diffusion_model
# remove from input blocks
_perform_adxl_motion_module_ejection(unet.input_blocks)
# remove from output blocks
_perform_adxl_motion_module_ejection(unet.output_blocks)
# remove from middle block (encapsulate in list to make structure compatible)
_perform_adxl_motion_module_ejection([unet.middle_block])
# remove attr; ejected
del_injected_unet_version(model)
def _perform_adxl_motion_module_ejection(unet_blocks: nn.ModuleList):
# eject all TemporalTransformer3DModel objects from all blocks
for block in unet_blocks:
idx_to_pop = []
for idx, component in enumerate(block):
if type(component) == VanillaTemporalModule:
idx_to_pop.append(idx)
# pop in backwards order, as to not disturb what the indeces refer to
for idx in sorted(idx_to_pop, reverse=True):
block.pop(idx)
#logger.info(f"ADXL: ejecting {idx_to_pop}")
############################################################################################################
############################################################################################################
## HotShot XL
def _inject_hsxl_motion_module_to_unet(model: ModelPatcher, motion_module: 'HotShotXLMotionWrapper'):
@@ -443,11 +529,13 @@ def _perform_hsxl_motion_module_ejection(unet_blocks: nn.ModuleList):
injectors = {
InjectorVersion.V1_V2: _inject_motion_module_to_unet,
InjectorVersion.ADXL_V1_V2: _inject_adxl_motion_module_to_unet,
InjectorVersion.HOTSHOTXL_V1: _inject_hsxl_motion_module_to_unet,
}
ejectors = {
InjectorVersion.V1_V2: _eject_motion_module_from_unet,
InjectorVersion.ADXL_V1_V2: _eject_adxl_motion_module_from_unet,
InjectorVersion.HOTSHOTXL_V1: _eject_hsxl_motion_module_from_unet,
}
+19 -2
View File
@@ -7,7 +7,7 @@ from torch import Tensor, nn
from comfy.ldm.modules.attention import FeedForward
from .motion_lora import MotionLoRAInfo
from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, TemporalTransformerGeneric, prepare_mask_batch
from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, MotionCompatibilityError, TemporalTransformerGeneric, prepare_mask_batch
def zero_module(module):
@@ -22,7 +22,23 @@ def get_ad_temporal_position_encoding_max_len(mm_state_dict: dict[str, Tensor],
for key in mm_state_dict.keys():
if key.endswith("pos_encoder.pe"):
return mm_state_dict[key].size(1) # get middle dim
raise ValueError(f"No pos_encoder.pe found in mm_state_dict - {mm_type} is not a valid motion module!")
raise MotionCompatibilityError(f"No pos_encoder.pe found in mm_state_dict - {mm_type} is not a valid AnimateDiff-SD1.5 motion module!")
def validate_ad_block_count(mm_state_dict: dict[str, Tensor], mm_type: str) -> None:
# keep track of biggest down_block count in module
biggest_block = 0
for key in mm_state_dict.keys():
if "down_blocks" in key:
try:
block_int = key.split(".")[1]
block_num = int(block_int)
if block_num > biggest_block:
biggest_block = block_num
except ValueError:
pass
if biggest_block != 3:
raise MotionCompatibilityError(f"Expected biggest down_block to be 3, but was {biggest_block} - {mm_type} is not a valid AnimateDiff-SD1.5 motion module!")
def has_mid_block(mm_state_dict: dict[str, Tensor]):
@@ -40,6 +56,7 @@ class AnimDiffMotionWrapper(GenericMotionWrapper):
self.up_blocks: Iterable[MotionModule] = nn.ModuleList([])
self.mid_block: Union[MotionModule, None] = None
self.encoding_max_len = get_ad_temporal_position_encoding_max_len(mm_state_dict, mm_name)
validate_ad_block_count(mm_state_dict, mm_name)
for c in (320, 640, 1280, 1280):
self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN))
for c in (1280, 1280, 640, 320):
+118
View File
@@ -0,0 +1,118 @@
import math
from typing import Callable, Iterable, Optional, Union
import torch
from einops import rearrange, repeat
from torch import Tensor, nn
from comfy.ldm.modules.attention import FeedForward
from .motion_lora import MotionLoRAInfo
from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, MotionCompatibilityError, TemporalTransformerGeneric, prepare_mask_batch
from .motion_module_ad import MotionModule
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
def get_ad_sdxl_temporal_position_encoding_max_len(mm_state_dict: dict[str, Tensor], mm_type: str) -> int:
# use pos_encoder.pe entries to determine max length - [1, {max_length}, {320|640|1280}]
for key in mm_state_dict.keys():
if key.endswith("pos_encoder.pe"):
return mm_state_dict[key].size(1) # get middle dim
raise MotionCompatibilityError(f"No pos_encoder.pe found in mm_state_dict - {mm_type} is not a valid AnimateDiff-SDXL motion module!")
def validate_ad_sdxl_block_count(mm_state_dict: dict[str, Tensor], mm_type: str) -> None:
# keep track of biggest down_block count in module
biggest_block = 0
for key in mm_state_dict.keys():
if "down_blocks" in key:
try:
block_int = key.split(".")[1]
block_num = int(block_int)
if block_num > biggest_block:
biggest_block = block_num
except ValueError:
pass
if biggest_block != 2:
raise MotionCompatibilityError(f"Expected biggest down_block to be 2, but was {biggest_block} - {mm_type} is not a valid AnimateDiff-SDXL motion module!")
def has_mid_block(mm_state_dict: dict[str, Tensor]):
# check if keys contain mid_block
for key in mm_state_dict.keys():
if key.startswith("mid_block."):
return True
return False
class AnimDiffSDXLMotionWrapper(GenericMotionWrapper):
def __init__(self, mm_state_dict: dict[str, Tensor], mm_hash: str, mm_name: str="mm_sd_v15.ckpt" , loras: list[MotionLoRAInfo]=None):
super().__init__(mm_hash, mm_name, loras)
self.down_blocks: Iterable[MotionModule] = nn.ModuleList([])
self.up_blocks: Iterable[MotionModule] = nn.ModuleList([])
self.mid_block: Union[MotionModule, None] = None
self.encoding_max_len = get_ad_sdxl_temporal_position_encoding_max_len(mm_state_dict, mm_name)
validate_ad_sdxl_block_count(mm_state_dict, mm_name)
for c in (320, 640, 1280):
self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN))
for c in (1280, 640, 320):
self.up_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.UP))
if has_mid_block(mm_state_dict):
self.mid_block = MotionModule(1280, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.MID)
self.mm_hash = mm_hash
self.mm_name = mm_name
self.version = "v1" if self.mid_block is None else "v2"
self.injector_version = InjectorVersion.ADXL_V1_V2
self.AD_video_length: int = 24
self.loras = loras
def has_loras(self):
# TODO: fix this to return False if has an empty list as well
# but only after implementing a fix for lowvram loading
return self.loras is not None
def set_video_length(self, video_length: int, full_length: int):
self.AD_video_length = video_length
for block in self.down_blocks:
block.set_video_length(video_length, full_length)
for block in self.up_blocks:
block.set_video_length(video_length, full_length)
if self.mid_block is not None:
self.mid_block.set_video_length(video_length, full_length)
def set_scale_multiplier(self, multiplier: Union[float, None]):
for block in self.down_blocks:
block.set_scale_multiplier(multiplier)
for block in self.up_blocks:
block.set_scale_multiplier(multiplier)
if self.mid_block is not None:
self.mid_block.set_scale_multiplier(multiplier)
def set_masks(self, masks: Tensor, min_val: float, max_val: float):
for block in self.down_blocks:
block.set_masks(masks, min_val, max_val)
for block in self.up_blocks:
block.set_masks(masks, min_val, max_val)
if self.mid_block is not None:
self.mid_block.set_masks(masks, min_val, max_val)
def set_sub_idxs(self, sub_idxs: list[int]):
for block in self.down_blocks:
block.set_sub_idxs(sub_idxs)
for block in self.up_blocks:
block.set_sub_idxs(sub_idxs)
if self.mid_block is not None:
self.mid_block.set_sub_idxs(sub_idxs)
def reset_temp_vars(self):
for block in self.down_blocks:
block.reset_temp_vars()
for block in self.up_blocks:
block.reset_temp_vars()
if self.mid_block is not None:
self.mid_block.reset_temp_vars()
+19 -2
View File
@@ -8,7 +8,7 @@ from torch import Tensor, nn
from comfy.ldm.modules.attention import FeedForward
from .motion_lora import MotionLoRAInfo
from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, TemporalTransformerGeneric
from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, MotionCompatibilityError, TemporalTransformerGeneric
def zero_module(module):
@@ -23,7 +23,23 @@ def get_hsxl_temporal_position_encoding_max_len(mm_state_dict: dict[str, Tensor]
for key in mm_state_dict.keys():
if key.endswith("pos_encoder.positional_encoding"):
return mm_state_dict[key].size(1) # get middle dim
raise ValueError(f"No pos_encoder.positional_encoding found in mm_state_dict - {mm_type} is not a valid HotShotXL motion module!")
raise MotionCompatibilityError(f"No pos_encoder.positional_encoding found in mm_state_dict - {mm_type} is not a valid HotShotXL motion module!")
def validate_hsxl_block_count(mm_state_dict: dict[str, Tensor], mm_type: str) -> None:
# keep track of biggest down_block count in module
biggest_block = 0
for key in mm_state_dict.keys():
if "down_blocks" in key:
try:
block_int = key.split(".")[1]
block_num = int(block_int)
if block_num > biggest_block:
biggest_block = block_num
except ValueError:
pass
if biggest_block != 2:
raise MotionCompatibilityError(f"Expected biggest down_block to be 2, but was {biggest_block} - {mm_type} is not a valid HotShotXL motion module!")
def has_mid_block(mm_state_dict: dict[str, Tensor]):
@@ -48,6 +64,7 @@ class HotShotXLMotionWrapper(GenericMotionWrapper):
self.up_blocks: Iterable[HotShotXLMotionModule] = nn.ModuleList([])
self.mid_block: Union[HotShotXLMotionModule, None] = None
self.encoding_max_len = get_hsxl_temporal_position_encoding_max_len(mm_state_dict, mm_name)
validate_hsxl_block_count(mm_state_dict, mm_name)
for c in (320, 640, 1280):
self.down_blocks.append(HotShotXLMotionModule(c, block_type=BlockType.DOWN, max_length=self.encoding_max_len))
for c in (1280, 640, 320):
+5
View File
@@ -139,6 +139,7 @@ class BlockType:
class InjectorVersion:
V1_V2 = "v1/v2"
ADXL_V1_V2 = "ADXL v1/v2"
HOTSHOTXL_V1 = "HSXL v1"
@@ -271,3 +272,7 @@ class NoiseType:
generator = torch.manual_seed(seed+i)
all_noises.append(torch.randn(single_shape, dtype=latents.dtype, layout=latents.layout, generator=generator, device="cpu"))
return torch.cat(all_noises, dim=0)
class MotionCompatibilityError(ValueError):
pass
+3 -2
View File
@@ -21,6 +21,7 @@ from .motion_module import InjectionParams, eject_motion_module, inject_motion_m
load_motion_module, unload_motion_module
from .motion_module import is_injected_mm_params, get_injected_mm_params
from .motion_module_ad import AnimDiffMotionWrapper, VanillaTemporalModule
from .motion_module_adxl import AnimDiffSDXLMotionWrapper
from .motion_utils import GenericMotionWrapper, GroupNormAD, NoiseType
@@ -155,8 +156,8 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable:
openaimodel.forward_timestep_embed = forward_timestep_embed
if params.unlimited_area_hack:
model_management.maximum_batch_area = unlimited_batch_area
# only apply groupnorm hack if not v2 and should not apply v2 properly
if not (isinstance(motion_module, AnimDiffMotionWrapper) and motion_module.version == "v2" and params.apply_v2_models_properly):
# only apply groupnorm hack if not [AnimateDiff and v2 and should apply v2 properly]
if not ((isinstance(motion_module, AnimDiffMotionWrapper) and motion_module.version == "v2" and params.apply_v2_models_properly)):
torch.nn.GroupNorm.forward = groupnorm_mm_factory(params)
if params.apply_mm_groupnorm_hack:
GroupNormAD.forward = groupnorm_mm_factory(params)