Fixed fp8 support by using comfy.ops
This commit is contained in:
@@ -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
|
||||
|
||||
+10
-7
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user