From a8a8f730af88388051c56a6b4ab71db2bb2e56b7 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 18 Jan 2024 10:58:40 -0600 Subject: [PATCH] Fixed fp8 support by using comfy.ops --- animatediff/motion_module_ad.py | 51 +++++++++++++++++++++------------ animatediff/sampling.py | 17 ++++++----- 2 files changed, 42 insertions(+), 26 deletions(-) diff --git a/animatediff/motion_module_ad.py b/animatediff/motion_module_ad.py index 5f9ab38..12c8089 100644 --- a/animatediff/motion_module_ad.py +++ b/animatediff/motion_module_ad.py @@ -12,9 +12,11 @@ from comfy.ldm.modules.diffusionmodules import openaimodel from comfy.ldm.modules.diffusionmodules.openaimodel import SpatialTransformer from comfy.controlnet import broadcast_image_to from comfy.utils import repeat_to_batch_size +import comfy.ops +import comfy.model_management 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_motion import CrossAttentionMM, MotionCompatibilityError, extend_to_batch_size, prepare_mask_batch from .utils_model import BetaSchedules, ModelTypeSD from .logger import logger @@ -148,7 +150,6 @@ class BlockType: MID = "mid" -# TODO: use comfy.ops style operations input for tensor operations to be in line with new ComfyUI practices class AnimateDiffModel(nn.Module): def __init__(self, mm_state_dict: dict[str, Tensor], mm_info: AnimateDiffInfo): super().__init__() @@ -157,6 +158,11 @@ class AnimateDiffModel(nn.Module): self.up_blocks: Iterable[MotionModule] = nn.ModuleList([]) self.mid_block: Union[MotionModule, None] = None self.encoding_max_len = get_position_encoding_max_len(mm_state_dict, mm_info.mm_name) + # determine ops to use (to support fp8 properly) + if comfy.model_management.unet_manual_cast(comfy.model_management.unet_dtype(), comfy.model_management.get_torch_device()) is None: + ops = comfy.ops.disable_weight_init + else: + ops = comfy.ops.manual_cast # SDXL has 3 up/down blocks, SD1.5 has 4 up/down blocks if mm_info.sd_type == ModelTypeSD.SDXL: layer_channels = (320, 640, 1280) @@ -164,11 +170,11 @@ class AnimateDiffModel(nn.Module): layer_channels = (320, 640, 1280, 1280) # fill out down/up blocks and middle block, if present for c in layer_channels: - self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN)) + self.down_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.DOWN, ops=ops)) for c in reversed(layer_channels): - self.up_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.UP)) + self.up_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.UP, ops=ops)) 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.mid_block = MotionModule(1280, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.MID, ops=ops) self.AD_video_length: int = 24 def get_device_debug(self): @@ -346,22 +352,22 @@ class AnimateDiffModel(nn.Module): class MotionModule(nn.Module): - def __init__(self, in_channels, temporal_position_encoding_max_len=24, block_type: str=BlockType.DOWN): + def __init__(self, in_channels, temporal_position_encoding_max_len=24, block_type: str=BlockType.DOWN, ops=comfy.ops.disable_weight_init): super().__init__() if block_type == BlockType.MID: # mid blocks contain only a single VanillaTemporalModule - self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList([get_motion_module(in_channels, temporal_position_encoding_max_len)]) + self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList([get_motion_module(in_channels, temporal_position_encoding_max_len, ops=ops)]) else: # down blocks contain two VanillaTemporalModules self.motion_modules: Iterable[VanillaTemporalModule] = nn.ModuleList( [ - get_motion_module(in_channels, temporal_position_encoding_max_len), - get_motion_module(in_channels, temporal_position_encoding_max_len) + get_motion_module(in_channels, temporal_position_encoding_max_len, ops=ops), + get_motion_module(in_channels, temporal_position_encoding_max_len, ops=ops) ] ) # up blocks contain one additional VanillaTemporalModule if block_type == BlockType.UP: - self.motion_modules.append(get_motion_module(in_channels, temporal_position_encoding_max_len)) + self.motion_modules.append(get_motion_module(in_channels, temporal_position_encoding_max_len, ops=ops)) def set_video_length(self, video_length: int, full_length: int): for motion_module in self.motion_modules: @@ -392,8 +398,8 @@ class MotionModule(nn.Module): motion_module.reset_temp_vars() -def get_motion_module(in_channels, temporal_position_encoding_max_len): - return VanillaTemporalModule(in_channels=in_channels, temporal_position_encoding_max_len=temporal_position_encoding_max_len) +def get_motion_module(in_channels, temporal_position_encoding_max_len, ops=comfy.ops.disable_weight_init): + return VanillaTemporalModule(in_channels=in_channels, temporal_position_encoding_max_len=temporal_position_encoding_max_len, ops=ops) class VanillaTemporalModule(nn.Module): @@ -408,6 +414,7 @@ class VanillaTemporalModule(nn.Module): temporal_position_encoding_max_len=24, temporal_attention_dim_div=1, zero_initialize=True, + ops=comfy.ops.disable_weight_init, ): super().__init__() @@ -431,6 +438,7 @@ class VanillaTemporalModule(nn.Module): cross_frame_attention_mode=cross_frame_attention_mode, temporal_position_encoding=temporal_position_encoding, temporal_position_encoding_max_len=temporal_position_encoding_max_len, + ops=ops ) if zero_initialize: @@ -528,6 +536,7 @@ class TemporalTransformer3DModel(nn.Module): cross_frame_attention_mode=None, temporal_position_encoding=False, temporal_position_encoding_max_len=24, + ops=comfy.ops.disable_weight_init, ): super().__init__() self.video_length = 16 @@ -540,10 +549,10 @@ class TemporalTransformer3DModel(nn.Module): inner_dim = num_attention_heads * attention_head_dim - self.norm = GroupNormAD( + self.norm = ops.GroupNorm( num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True ) - self.proj_in = nn.Linear(in_channels, inner_dim) + self.proj_in = ops.Linear(in_channels, inner_dim) self.transformer_blocks: Iterable[TemporalTransformerBlock] = nn.ModuleList( [ @@ -561,11 +570,12 @@ class TemporalTransformer3DModel(nn.Module): cross_frame_attention_mode=cross_frame_attention_mode, temporal_position_encoding=temporal_position_encoding, temporal_position_encoding_max_len=temporal_position_encoding_max_len, + ops=ops, ) for d in range(num_layers) ] ) - self.proj_out = nn.Linear(inner_dim, in_channels) + self.proj_out = ops.Linear(inner_dim, in_channels) def set_video_length(self, video_length: int, full_length: int): self.video_length = video_length @@ -687,6 +697,7 @@ class TemporalTransformerBlock(nn.Module): cross_frame_attention_mode=None, temporal_position_encoding=False, temporal_position_encoding_max_len=24, + ops=comfy.ops.disable_weight_init, ): super().__init__() @@ -709,15 +720,16 @@ class TemporalTransformerBlock(nn.Module): cross_frame_attention_mode=cross_frame_attention_mode, temporal_position_encoding=temporal_position_encoding, temporal_position_encoding_max_len=temporal_position_encoding_max_len, + ops=ops, ) ) - norms.append(nn.LayerNorm(dim)) + norms.append(ops.LayerNorm(dim)) self.attention_blocks: Iterable[VersatileAttention] = nn.ModuleList(attention_blocks) self.norms = nn.ModuleList(norms) - self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu")) - self.ff_norm = nn.LayerNorm(dim) + self.ff = FeedForward(dim, dropout=dropout, glu=(activation_fn == "geglu"), operations=ops) + self.ff_norm = ops.LayerNorm(dim) def set_scale_multiplier(self, multiplier: Union[float, None]): for block in self.attention_blocks: @@ -826,10 +838,11 @@ class VersatileAttention(CrossAttentionMM): cross_frame_attention_mode=None, temporal_position_encoding=False, temporal_position_encoding_max_len=24, + ops=comfy.ops.disable_weight_init, *args, **kwargs, ): - super().__init__(*args, **kwargs) + super().__init__(operations=ops, *args, **kwargs) assert attention_mode == "Temporal" self.attention_mode = attention_mode diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 342a2b5..ee09175 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -13,10 +13,10 @@ import comfy.samplers import comfy.sample import comfy.utils from comfy.controlnet import ControlBase +import comfy.ops from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows from .sample_settings import IterationOptions, SeedNoiseGeneration, prepare_mask_ad -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 @@ -112,7 +112,7 @@ def unlimited_memory_required(*args, **kwargs): return 0 -def groupnorm_mm_factory(params: InjectionParams): +def groupnorm_mm_factory(params: InjectionParams, manual_cast=False): def groupnorm_mm_forward(self, input: Tensor) -> Tensor: # 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 @@ -122,7 +122,11 @@ def groupnorm_mm_factory(params: InjectionParams): batched_conds = input.size(0)//params.context_options.context_length 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) + if manual_cast: + weight, bias = comfy.ops.cast_bias_weight(self, input) + else: + weight, bias = self.weight, self.bias + input = group_norm(input, self.num_groups, weight, bias, self.eps) input = rearrange(input, "b c f h w -> (b f) c h w", b=batched_conds) return input return groupnorm_mm_forward @@ -189,7 +193,7 @@ class FunctionInjectionHolder: self.orig_forward_timestep_embed = openaimodel.forward_timestep_embed # needed to account for VanillaTemporalModule self.orig_memory_required = model.model.memory_required # allows for "unlimited area hack" to prevent halving of conds/unconds self.orig_groupnorm_forward = torch.nn.GroupNorm.forward # used to normalize latents to remove "flickering" of colors/brightness between frames - self.orig_groupnormad_forward = GroupNormAD.forward + self.orig_groupnorm_manual_cast_forward = comfy.ops.manual_cast.GroupNorm.forward_comfy_cast_weights self.orig_sampling_function = comfy.samplers.sampling_function # used to support sliding context windows in samplers self.orig_prepare_mask = comfy.sample.prepare_mask self.orig_get_additional_models = comfy.sample.get_additional_models @@ -203,8 +207,7 @@ class FunctionInjectionHolder: if not (info.mm_version == AnimateDiffVersion.V3 or (info.mm_format == AnimateDiffFormat.ANIMATEDIFF and info.sd_type == ModelTypeSD.SD1_5 and info.mm_version == AnimateDiffVersion.V2 and params.apply_v2_properly)): torch.nn.GroupNorm.forward = groupnorm_mm_factory(params) - if params.apply_mm_groupnorm_hack: - GroupNormAD.forward = groupnorm_mm_factory(params) + comfy.ops.manual_cast.GroupNorm.forward_comfy_cast_weights = groupnorm_mm_factory(params, manual_cast=True) # if mps device (Apple Silicon), disable batched conds to avoid black images with groupnorm hack try: if model.load_device.type == "mps": @@ -222,7 +225,7 @@ class FunctionInjectionHolder: model.model.memory_required = self.orig_memory_required openaimodel.forward_timestep_embed = self.orig_forward_timestep_embed torch.nn.GroupNorm.forward = self.orig_groupnorm_forward - GroupNormAD.forward = self.orig_groupnormad_forward + comfy.ops.manual_cast.GroupNorm.forward_comfy_cast_weights = self.orig_groupnorm_manual_cast_forward comfy.samplers.sampling_function = self.orig_sampling_function comfy.sample.prepare_mask = self.orig_prepare_mask comfy.sample.get_additional_models = self.orig_get_additional_models