diff --git a/animatediff/model_utils.py b/animatediff/model_utils.py index 9395a15..c9fe398 100644 --- a/animatediff/model_utils.py +++ b/animatediff/model_utils.py @@ -3,8 +3,6 @@ import hashlib import folder_paths -from .logger import logger - folder_paths.folder_names_and_paths["AnimateDiff"] = ( [ @@ -14,17 +12,6 @@ folder_paths.folder_names_and_paths["AnimateDiff"] = ( folder_paths.supported_pt_extensions, ) -known_models = { - "aa7fd8a200a89031edd84487e2a757c5315460eca528fa70d4b3885c399bffd5": "mm_sd_v14.ckpt", - "cf16ea656cb16124990c8e2c70a29c793f9841f3a2223073fac8bd89ebd9b69a": "mm_sd_v15.ckpt", - "0aaf157b9c51a0ae07cb5d9ea7c51299f07bddc6f52025e1f9bb81cd763631df": "mm-Stabilized_high.pth", - "39de8b71b1c09f10f4602f5d585d82771a60d3cf282ba90215993e06afdfe875": "mm-Stabilized_mid.pth", - "3cb569f7ce3dc6a10aa8438e666265cb9be3120d8f205de6a456acf46b6c99f4": "temporaldiff-v1-animatediff.ckpt", - "69ed0f5fef82b110aca51bcab73b21104242bc65d6ab4b8b2a2a94d31cad1bf0": "mm_sd_v15_v2.ckpt", -} - -v2_models = ["69ed0f5fef82b110aca51bcab73b21104242bc65d6ab4b8b2a2a94d31cad1bf0"] - def get_available_models(): return folder_paths.get_filename_list("AnimateDiff") @@ -34,25 +21,7 @@ def get_model_path(model_name): return folder_paths.get_full_path("AnimateDiff", model_name) -def sha256_file(file_path): +def get_model_hash(file_path): with open(file_path, "rb") as f: bytes = f.read() # read entire file as bytes return hashlib.sha256(bytes).hexdigest() - - -def validate_mm_model(model_name): - model_path = get_model_path(model_name) - model_hash = sha256_file(model_path) - - if model_hash in known_models: - logger.info(f"You are using {model_name}, which has been tested and supported.") - else: - logger.warn( - f"Your model {model_name} has not been tested and supported." - "Either your download is incomplete or your model has not been tested. " - "Please use at your own risk." - ) - - using_v2 = model_hash in v2_models - - return (model_hash, using_v2) \ No newline at end of file diff --git a/animatediff/motion_module.py b/animatediff/motion_module.py index 300829c..4ed20bd 100644 --- a/animatediff/motion_module.py +++ b/animatediff/motion_module.py @@ -1,12 +1,12 @@ +import os import torch -import torch.nn.functional as F -from torch import nn +from torch import Tensor, nn import math from einops import rearrange, repeat -from comfy.ldm.modules.attention import FeedForward -from .attention_processor import Attention as CrossAttention +from comfy.utils import load_torch_file +from comfy.ldm.modules.attention import FeedForward, CrossAttention def zero_module(module): @@ -15,41 +15,108 @@ def zero_module(module): p.detach().zero_() return module + +# Merge from https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved +def get_encoding_max_len(mm_state_dict: dict[str, Tensor]) -> 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 ValueError(f"No pos_encoder.pe found in mm_state_dict") + + +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 MotionWrapper(nn.Module): - def __init__(self, mm_hash, is_v2 = False): + def __init__(self, mm_type: str, encoding_max_len: int = 24, is_v2=False): super().__init__() - if is_v2: - max_len = 32 - else: - max_len = 24 + self.mm_type = mm_type + self.is_v2 = is_v2 self.down_blocks = nn.ModuleList([]) self.up_blocks = nn.ModuleList([]) + self.mid_block = None + for c in (320, 640, 1280, 1280): - self.down_blocks.append(MotionModule(c, max_len=max_len)) + self.down_blocks.append( + MotionModule(c, BlockType.DOWN, encoding_max_len=encoding_max_len) + ) for c in (1280, 1280, 640, 320): - self.up_blocks.append(MotionModule(c, is_up=True, max_len=max_len)) + self.up_blocks.append( + MotionModule(c, BlockType.UP, encoding_max_len=encoding_max_len) + ) if is_v2: - self.mid_block = MotionModule(1280, max_len=max_len, is_mid=is_v2) - self.mm_hash = mm_hash - self.is_v2 = is_v2 + self.mid_block = MotionModule( + 1280, BlockType.MID, encoding_max_len=encoding_max_len + ) + + @classmethod + def from_pretrained(cls, checkpoint_path: str): + mm_state_dict = load_torch_file(checkpoint_path) + mm_type = os.path.basename(checkpoint_path) + encoding_max_len = get_encoding_max_len(mm_state_dict) + is_v2 = has_mid_block(mm_state_dict) + + mm = cls(mm_type, encoding_max_len=encoding_max_len, is_v2=is_v2) + mm.load_state_dict(mm_state_dict) + return mm + + def set_video_length(self, video_length: int): + for block in self.down_blocks: + block.set_video_length(video_length) + for block in self.up_blocks: + block.set_video_length(video_length) + if self.mid_block is not None: + self.mid_block.set_video_length(video_length) + + +class BlockType: + UP = "up" + DOWN = "down" + MID = "mid" class MotionModule(nn.Module): - def __init__(self, in_channels, is_up=False, is_mid=False, max_len=24): + def __init__( + self, + in_channels, + block_type: BlockType, + encoding_max_len=24, + ): super().__init__() - if is_mid: - self.motion_modules = nn.ModuleList([get_motion_module(in_channels, max_len)]) + self.block_type = block_type + + if block_type == BlockType.MID: + self.motion_modules = nn.ModuleList( + [get_motion_module(in_channels, encoding_max_len)] + ) else: self.motion_modules = nn.ModuleList( - [get_motion_module(in_channels, max_len), get_motion_module(in_channels, max_len)] + [ + get_motion_module(in_channels, encoding_max_len), + get_motion_module(in_channels, encoding_max_len), + ] ) - if is_up: - self.motion_modules.append(get_motion_module(in_channels, max_len)) + if block_type == BlockType.UP: + self.motion_modules.append( + get_motion_module(in_channels, encoding_max_len) + ) + + def set_video_length(self, video_length: int): + for motion_module in self.motion_modules: + motion_module.set_video_length(video_length) def get_motion_module(in_channels, max_len): - return VanillaTemporalModule(in_channels=in_channels, temporal_position_encoding_max_len=max_len) + return VanillaTemporalModule( + in_channels=in_channels, temporal_position_encoding_max_len=max_len + ) class VanillaTemporalModule(nn.Module): @@ -85,8 +152,13 @@ class VanillaTemporalModule(nn.Module): self.temporal_transformer.proj_out ) + def set_video_length(self, video_length: int): + self.temporal_transformer.set_video_length(video_length) + def forward(self, input_tensor, encoder_hidden_states, attention_mask=None): - return self.temporal_transformer(input_tensor, encoder_hidden_states, attention_mask) + return self.temporal_transformer( + input_tensor, encoder_hidden_states, attention_mask + ) class TemporalTransformer3DModel(nn.Module): @@ -140,10 +212,12 @@ class TemporalTransformer3DModel(nn.Module): ] ) self.proj_out = nn.Linear(inner_dim, in_channels) + self.video_length = 16 + + def set_video_length(self, video_length: int): + self.video_length = video_length def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None): - video_length = hidden_states.shape[0] // 2 # TODO: config this value in scripts - batch, channel, height, weight = hidden_states.shape residual = hidden_states @@ -159,7 +233,7 @@ class TemporalTransformer3DModel(nn.Module): hidden_states = block( hidden_states, encoder_hidden_states=encoder_hidden_states, - video_length=video_length, + video_length=self.video_length, ) # output @@ -204,15 +278,15 @@ class TemporalTransformerBlock(nn.Module): attention_blocks.append( VersatileAttention( attention_mode=block_name.split("_")[0], - cross_attention_dim=cross_attention_dim + context_dim=cross_attention_dim if block_name.endswith("_Cross") else None, query_dim=dim, heads=num_attention_heads, dim_head=attention_head_dim, dropout=dropout, - bias=attention_bias, - upcast_attention=upcast_attention, + # bias=attention_bias, # remove for Comfy CrossAttention + # upcast_attention=upcast_attention, # remove for Comfy CrossAttention cross_frame_attention_mode=cross_frame_attention_mode, temporal_position_encoding=temporal_position_encoding, temporal_position_encoding_max_len=temporal_position_encoding_max_len, @@ -284,7 +358,7 @@ class VersatileAttention(CrossAttention): assert attention_mode == "Temporal" self.attention_mode = attention_mode - self.is_cross_attention = kwargs["cross_attention_dim"] is not None + self.is_cross_attention = kwargs["context_dim"] is not None self.pos_encoder = ( PositionalEncoding( @@ -327,8 +401,8 @@ class VersatileAttention(CrossAttention): hidden_states = super().forward( hidden_states, encoder_hidden_states, - attention_mask, - **cross_attention_kwargs, + value=None, + mask=attention_mask, ) hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d) diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 1b92303..1ff8191 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -1,9 +1,10 @@ import os import json -import hashlib import torch import numpy as np -from typing import Dict, List, Tuple +from typing import Dict, List +from torch import Tensor +from torch.nn.functional import group_norm from PIL import Image from PIL.PngImagePlugin import PngInfo from einops import rearrange @@ -11,19 +12,14 @@ from einops import rearrange import folder_paths import comfy.ldm.modules.diffusionmodules.openaimodel as openaimodel import comfy.model_management as model_management +from comfy.model_base import BaseModel from comfy.ldm.modules.attention import SpatialTransformer -from comfy.ldm.modules.diffusionmodules.util import GroupNorm32 -from comfy.utils import load_torch_file, calculate_parameters -from comfy.model_patcher import ModelPatcher +from comfy.cli_args import args as cli_args from nodes import KSampler from .logger import logger from .motion_module import MotionWrapper, VanillaTemporalModule -from .model_utils import get_available_models, get_model_path, validate_mm_model - - -orig_forward_timestep_embed = openaimodel.forward_timestep_embed -groupnorm32_original_forward = GroupNorm32.forward +from .model_utils import get_available_models, get_model_path, get_model_hash def forward_timestep_embed( @@ -44,44 +40,38 @@ def forward_timestep_embed( return x -def groupnorm32_mm_forward(self, x): - x = rearrange(x, "(b f) c h w -> b c f h w", b=2) - x = groupnorm32_original_forward(self, x) - x = rearrange(x, "b c f h w -> (b f) c h w", b=2) - return x +def groupnorm_mm_factory(video_length: int): + 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 + axes_factor = input.size(0) // video_length + + input = rearrange(input, "(b f) c h w -> b c f h w", b=axes_factor) + input = group_norm(input, self.num_groups, self.weight, self.bias, self.eps) + input = rearrange(input, "b c f h w -> (b f) c h w", b=axes_factor) + return input + + return groupnorm_mm_forward +orig_forward_timestep_embed = openaimodel.forward_timestep_embed +orig_maximum_batch_area = model_management.maximum_batch_area +orig_groupnorm_forward = torch.nn.GroupNorm.forward openaimodel.forward_timestep_embed = forward_timestep_embed motion_modules: Dict[str, MotionWrapper] = {} -original_model_hashs = set() -injected_model_hashs: Dict[str, Tuple[str, str]] = {} - - -def calculate_model_hash(unet): - t = unet.input_blocks[1] - m = hashlib.sha256() - for buf in t.buffers(): - m.update(buf.cpu().numpy().view(np.uint8)) - return m.hexdigest() def load_motion_module(model_name: str): model_path = get_model_path(model_name) - model_hash, is_v2 = validate_mm_model(model_name) + model_hash = get_model_hash(model_path) if model_hash not in motion_modules: logger.info(f"Loading motion module {model_name}") - mm_state_dict = load_torch_file(model_path) - motion_module = MotionWrapper(model_name, is_v2=is_v2) - - parameters = calculate_parameters(mm_state_dict, "") - usefp16 = model_management.should_use_fp16(model_params=parameters) - if usefp16: - logger.info("Using fp16, converting motion module to fp16") + motion_module = MotionWrapper.from_pretrained(model_path) + if not cli_args.force_fp32: + logger.info(f"Converting motion module to fp16.") motion_module.half() - # offload_device = model_management.unet_offload_device() - # motion_module = motion_module.to(offload_device) - motion_module.load_state_dict(mm_state_dict) + motion_modules[model_hash] = motion_module return motion_modules[model_hash] @@ -176,81 +166,6 @@ ejectors = { } -class AnimateDiffLoaderLegacy: - def __init__(self) -> None: - self.version = "legacy" - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "model_name": (get_available_models(),), - "width": ("INT", {"default": 512, "min": 64, "max": 1024, "step": 8}), - "height": ("INT", {"default": 512, "min": 64, "max": 1024, "step": 8}), - "frame_number": ( - "INT", - {"default": 16, "min": 2, "max": 24, "step": 1}, - ), - }, - "optional": { - "init_latent": ("LATENT",), - }, - } - - @classmethod - def IS_CHANGED(s, model: ModelPatcher): - unet = model.model.diffusion_model - # return calculate_model_hash(unet) not in injected_model_hashs - return hasattr(unet, "motion_module") and unet.motion_module is not None - - RETURN_TYPES = ("MODEL", "LATENT") - CATEGORY = "Animate Diff" - FUNCTION = "inject_motion_modules" - - def inject_motion_modules( - self, - model: ModelPatcher, - model_name: str, - width: int, - height: int, - frame_number=16, - init_latent: Dict[str, torch.Tensor] = None, - ): - motion_module = load_motion_module(model_name) - - model = model.clone() - unet = model.model.diffusion_model - unet_hash = calculate_model_hash(unet) - need_inject = unet_hash not in injected_model_hashs - - if unet_hash in injected_model_hashs: - (mm_hash, version) = injected_model_hashs[unet_hash] - if version != self.version or mm_hash != motion_module.mm_hash: - # injected by another motion module, unload first - logger.info(f"Ejecting motion module {mm_hash} version {version}.") - ejectors[version](unet) - need_inject = True - else: - logger.info(f"Motion module already injected, skipping injection.") - - if need_inject: - logger.info(f"Injecting motion module {model_name} version {self.version}.") - injectors[self.version](unet, motion_module) - unet_hash = calculate_model_hash(unet) - injected_model_hashs[unet_hash] = (motion_module.mm_hash, self.version) - - if init_latent is None: - latent = torch.zeros([frame_number, 4, height // 8, width // 8]).cpu() - else: - # clone value of first frame - latent = init_latent["samples"][:1, :, :, :].clone().cpu() - # repeat for all frames - latent = latent.repeat(frame_number, 1, 1, 1) - - return (model, {"samples": latent}) - - class AnimateDiffModuleLoader: @classmethod def INPUT_TYPES(s): @@ -273,105 +188,6 @@ class AnimateDiffModuleLoader: return (motion_module,) -class AnimateDiffLoader: - def __init__(self) -> None: - self.version = "v1" - - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "init_latent": ("LATENT",), - "model_name": (get_available_models(),), - "frame_number": ( - "INT", - {"default": 16, "min": 2, "max": 32, "step": 1}, - ), - }, - } - - @classmethod - def IS_CHANGED(s, model: ModelPatcher, _): - unet = model.model.diffusion_model - # return calculate_model_hash(unet) not in injected_model_hashs - return hasattr(unet, "motion_module") and unet.motion_module is not None - - RETURN_TYPES = ("MODEL", "LATENT") - CATEGORY = "Animate Diff" - FUNCTION = "inject_motion_modules" - - def inject_motion_modules( - self, - model: ModelPatcher, - init_latent: Dict[str, torch.Tensor], - model_name: str, - frame_number=16, - ): - motion_module = load_motion_module(model_name) - - model = model.clone() - unet = model.model.diffusion_model - unet_hash = calculate_model_hash(unet) - need_inject = unet_hash not in injected_model_hashs - - if unet_hash in injected_model_hashs: - (mm_type, version) = injected_model_hashs[unet_hash] - if version != self.version or mm_type != motion_module.mm_hash: - # injected by another motion module, unload first - logger.info(f"Ejecting motion module {mm_type} version {version}.") - ejectors[version](unet) - need_inject = True - else: - logger.info(f"Motion module already injected, skipping injection.") - - if need_inject: - logger.info(f"Injecting motion module {model_name} version {self.version}.") - injectors[self.version](unet, motion_module) - unet_hash = calculate_model_hash(unet) - injected_model_hashs[unet_hash] = (motion_module.mm_hash, self.version) - - init_frames = len(init_latent["samples"]) - samples = init_latent["samples"][:init_frames, :, :, :].clone().cpu() - - if init_frames < frame_number: - last_frame = samples[-1].unsqueeze(0) - repeated_last_frames = last_frame.repeat( - frame_number - init_frames, 1, 1, 1 - ) - samples = torch.cat((samples, repeated_last_frames), dim=0) - - return (model, {"samples": samples}) - - -class AnimateDiffUnload: - @classmethod - def INPUT_TYPES(s): - return {"required": {"model": ("MODEL",)}} - - @classmethod - def IS_CHANGED(s, model: ModelPatcher): - unet = model.model.diffusion_model - return calculate_model_hash(unet) in injected_model_hashs - - RETURN_TYPES = ("MODEL",) - CATEGORY = "Animate Diff" - FUNCTION = "unload_motion_modules" - - def unload_motion_modules(self, model: ModelPatcher): - model = model.clone() - unet = model.model.diffusion_model - model_hash = calculate_model_hash(unet) - if model_hash in injected_model_hashs: - (model_name, version) = injected_model_hashs[model_hash] - logger.info(f"Ejecting motion module {model_name} version {version}.") - ejectors[version](unet) - else: - logger.info(f"Motion module not injected, skip unloading.") - - return (model,) - - class AnimateDiffSampler(KSampler): @classmethod def INPUT_TYPES(s): @@ -394,66 +210,56 @@ class AnimateDiffSampler(KSampler): def __init__(self) -> None: super().__init__() self.prev_beta = None - self.prev_alpha_cumprod = None - self.prev_alpha_cumprod_prev = None + self.prev_linear_start = None + self.prev_linear_end = None - def override_ddim_alpha(self, model): - logger.info(f"Setting DDIM alpha.") - device = model_management.unet_offload_device() - - beta_start = 0.00085 - beta_end = 0.012 - betas = torch.linspace( - beta_start, - beta_end, - model.num_timesteps, - dtype=torch.float32, - device=device, + def override_beta_schedule(self, model: BaseModel): + logger.info(f"Override beta schedule.") + self.prev_beta = model.get_buffer("betas") + self.prev_linear_start = model.linear_start + self.prev_linear_end = model.linear_end + model.register_schedule( + given_betas=None, + beta_schedule="sqrt_linear", + timesteps=1000, + linear_start=0.00085, + linear_end=0.012, + cosine_s=8e-3, ) - alphas = 1.0 - betas - alphas_cumprod = torch.cumprod(alphas, dim=0) - alphas_cumprod_prev = torch.cat( - ( - torch.tensor([1.0], dtype=torch.float32, device=device), - alphas_cumprod[:-1], - ) - ) - self.prev_beta = model.betas - model.betas = betas - self.prev_alpha_cumprod = model.alphas_cumprod - model.alphas_cumprod = alphas_cumprod - self.prev_alpha_cumprod_prev = model.alphas_cumprod_prev - model.alphas_cumprod_prev = alphas_cumprod_prev - def restore_ddim_alpha(self, model): - logger.info(f"Restoring DDIM alpha.") - model.betas = self.prev_beta - model.alphas_cumprod = self.prev_alpha_cumprod - model.alphas_cumprod_prev = self.prev_alpha_cumprod_prev + def restore_beta_schedule(self, model: BaseModel): + logger.info(f"Restoring beta schedule.") + model.register_schedule( + given_betas=self.prev_beta, + linear_start=self.prev_linear_start, + linear_end=self.prev_linear_end, + ) self.prev_beta = None - self.prev_alpha_cumprod = None - self.prev_alpha_cumprod_prev = None + self.prev_linear_start = None + self.prev_linear_end = None - def inject_motion_module(self, model, motion_module, inject_method): + def inject_motion_module( + self, model, motion_module: MotionWrapper, inject_method: str, frame_number: int + ): model = model.clone() unet = model.model.diffusion_model logger.info(f"Injecting motion module with method {inject_method}.") injectors[inject_method](unet, motion_module) - self.override_ddim_alpha(model.model) + self.override_beta_schedule(model.model) if not motion_module.is_v2: - logger.info(f"Hacking GroupNorm32 forward function.") - GroupNorm32.forward = groupnorm32_mm_forward + logger.info(f"Hacking GroupNorm.forward function.") + torch.nn.GroupNorm.forward = groupnorm_mm_factory(frame_number) return model def eject_motion_module(self, model, inject_method): unet = model.model.diffusion_model - self.restore_ddim_alpha(model.model) + self.restore_beta_schedule(model.model) if not unet.motion_module.is_v2: logger.info(f"Restore GroupNorm32 forward function.") - GroupNorm32.forward = groupnorm32_original_forward + torch.nn.GroupNorm.forward = orig_groupnorm_forward logger.info(f"Ejecting motion module with method {inject_method}.") ejectors[inject_method](unet) @@ -474,7 +280,9 @@ class AnimateDiffSampler(KSampler): latent_image, denoise=1.0, ): - model = self.inject_motion_module(model, motion_module, inject_method) + model = self.inject_motion_module( + model, motion_module, inject_method, frame_number + ) init_frames = len(latent_image["samples"]) samples = latent_image["samples"][:init_frames, :, :, :].clone().cpu() @@ -488,22 +296,23 @@ class AnimateDiffSampler(KSampler): latent_image = {"samples": samples} - results = super().sample( - model, - seed, - steps, - cfg, - sampler_name, - scheduler, - positive, - negative, - latent_image, - denoise=1.0, - ) - - self.eject_motion_module(model, inject_method) - - return results + try: + return super().sample( + model, + seed, + steps, + cfg, + sampler_name, + scheduler, + positive, + negative, + latent_image, + denoise=denoise, + ) + except: + raise + finally: + self.eject_motion_module(model, inject_method) class AnimateDiffCombine: @@ -603,17 +412,11 @@ class AnimateDiffCombine: NODE_CLASS_MAPPINGS = { - # "AnimateDiffLoader": AnimateDiffLoaderLegacy, - # "AnimateDiffLoader_v2": AnimateDiffLoader, - # "AnimateDiffUnload": AnimateDiffUnload, "AnimateDiffModuleLoader": AnimateDiffModuleLoader, "AnimateDiffCombine": AnimateDiffCombine, "AnimateDiffSampler": AnimateDiffSampler, } NODE_DISPLAY_NAME_MAPPINGS = { - # "AnimateDiffLoader": "[DEPRECATED] Animate Diff Loader Legacy", - # "AnimateDiffLoader_v2": "[DEPRECATED] Animate Diff Loader", - # "AnimateDiffUnload": "[DEPRECATED] Animate Diff Unload", "AnimateDiffModuleLoader": "Animate Diff Module Loader", "AnimateDiffSampler": "Animate Diff Sampler", "AnimateDiffCombine": "Animate Diff Combine",