Fixed fp8 support by using comfy.ops

This commit is contained in:
Jedrzej Kosinski
2024-01-18 10:58:40 -06:00
parent b459e3e0ef
commit a8a8f730af
2 changed files with 42 additions and 26 deletions
+32 -19
View File
@@ -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
View File
@@ -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