Added AnimateDiff-SDXL support - you need to use 'linear (AnimateDiff-SDXL)' beta_schedule
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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()
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user