From 63c068c59fd385e9c5b6fbb23020d963062835a9 Mon Sep 17 00:00:00 2001 From: "Tung Nguyen (Blockchain)" <91741875+tungnguyensipher@users.noreply.github.com> Date: Mon, 25 Sep 2023 22:53:44 +0700 Subject: [PATCH] Support motion LoRA (#38) Support new motion LoRa from AnimateDiff --- README.md | 42 +++ animatediff/model_utils.py | 46 +++- animatediff/nodes.py | 86 ++++++- loras/.gitkeep | 0 workflows/lora.json | 515 +++++++++++++++++++++++++++++++++++++ 5 files changed, 686 insertions(+), 3 deletions(-) create mode 100644 loras/.gitkeep create mode 100644 workflows/lora.json diff --git a/README.md b/README.md index b70b2fc..a641544 100644 --- a/README.md +++ b/README.md @@ -11,6 +11,48 @@ - Community modules: [manshoety/AD_Stabilized_Motion](https://huggingface.co/manshoety/AD_Stabilized_Motion) | [CiaraRowles/TemporalDiff](https://huggingface.co/CiaraRowles/TemporalDiff) - AnimateDiff v2 [mm_sd_v15_v2.ckpt](https://huggingface.co/guoyww/animatediff/blob/main/mm_sd_v15_v2.ckpt) +## Update 2023/09/25 + +#### **Motion LoRA** is now supported! + +Download [motion LoRAs](https://huggingface.co/guoyww/animatediff/tree/main) and put them under `comfyui-animatediff/loras/` folder. + +Note: LoRAs only work with **AnimateDiff v2** [mm_sd_v15_v2.ckpt](https://huggingface.co/guoyww/animatediff/blob/main/mm_sd_v15_v2.ckpt) module. + +#### New node: `AnimateDiffLoraLoader` + +image + +Example workflow: +image + +Workflow: [lora.json](https://github.com/ArtVentureX/comfyui-animatediff/blob/main/workflows/lora.json) + +Samples: + + + + + + + + + + + + + + +
+image +
+image +
+image +
+image +
+ ## Update 2023/09/21 #### **Sliding Window** is now available! diff --git a/animatediff/model_utils.py b/animatediff/model_utils.py index 688a3f4..321f6c4 100644 --- a/animatediff/model_utils.py +++ b/animatediff/model_utils.py @@ -1,5 +1,6 @@ import os import hashlib +import torch from typing import Dict import folder_paths @@ -11,6 +12,7 @@ from .motion_module import MotionWrapper motion_modules: Dict[str, MotionWrapper] = {} +motion_loras: Dict[str, Dict[str, torch.Tensor]] = {} folder_paths.folder_names_and_paths["AnimateDiff"] = ( @@ -20,19 +22,34 @@ folder_paths.folder_names_and_paths["AnimateDiff"] = ( ], folder_paths.supported_pt_extensions, ) +folder_paths.folder_names_and_paths["AnimateDiffLora"] = ( + [ + os.path.join(folder_paths.models_dir, "AnimateDiffLora"), + os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "loras"), + ], + folder_paths.supported_pt_extensions, +) def get_available_models(): return folder_paths.get_filename_list("AnimateDiff") +def get_available_loras(): + return folder_paths.get_filename_list("AnimateDiffLora") + + def get_model_path(model_name): return folder_paths.get_full_path("AnimateDiff", model_name) +def get_lora_path(lora_name): + return folder_paths.get_full_path("AnimateDiffLora", lora_name) + + def get_model_hash(file_path): with open(file_path, "rb") as f: - bytes = f.read() # read entire file as bytes + bytes = f.read(1024 * 1024) # read entire file as bytes return hashlib.sha256(bytes).hexdigest() @@ -54,3 +71,30 @@ def load_motion_module(model_name: str): motion_modules[model_hash] = motion_module return motion_modules[model_hash] + + +def load_lora(lora_name: str): + lora_path = get_lora_path(lora_name) + lora_hash = get_model_hash(lora_path) + if lora_hash not in motion_modules: + logger.info(f"Loading lora {lora_name}") + state_dict = load_torch_file(lora_path) + updated_state_dict: Dict[str, torch.Tensor] = {} + + for key in state_dict: + # only process lora down key + if "up." in key: + continue + + up_key = key.replace(".down.", ".up.") + model_key = key.replace("processor.", "").replace("_lora", "").replace("down.", "").replace("up.", "") + model_key = model_key.replace("to_out.", "to_out.0.") + combined_key = ".".join(model_key.split(".")[:-1]) + + weight_down = state_dict[key] + weight_up = state_dict[up_key] + updated_state_dict[combined_key] = torch.mm(weight_up, weight_down).to("cpu") + + motion_loras[lora_hash] = updated_state_dict + + return motion_loras[lora_hash] diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 3af01ce..9b06dc9 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -3,16 +3,18 @@ import json import torch import numpy as np import hashlib -from typing import List +from typing import List, Dict, Tuple from torch import Tensor from PIL import Image, ImageSequence from PIL.PngImagePlugin import PngInfo import folder_paths -from .model_utils import get_available_models, load_motion_module +from .motion_module import MotionWrapper +from .model_utils import get_available_models, load_motion_module, get_available_loras, load_lora from .utils import pil2tensor, ensure_opencv from .sampler import AnimateDiffSampler, AnimateDiffSlidingWindowOptions +from .logger import logger SLIDING_CONTEXT_LENGTH = 16 @@ -28,21 +30,99 @@ class AnimateDiffModuleLoader: "required": { "model_name": (get_available_models(),), }, + "optional": { + "lora_stack": ("MOTION_LORA_STACK",), + }, } RETURN_TYPES = ("MOTION_MODULE",) CATEGORY = "Animate Diff" FUNCTION = "load_motion_module" + def inject_loras(self, motion_module: MotionWrapper, lora_stack: List[Tuple[Dict[str, Tensor], float]]): + for lora in lora_stack: + (state_dict, alpha) = lora + + for key in state_dict: + layer_infos = key.split(".") + + curr_layer = motion_module + while len(layer_infos) > 0: + temp_name = layer_infos.pop(0) + curr_layer = curr_layer.__getattr__(temp_name) + + curr_layer.weight.data += alpha * state_dict[key].to(curr_layer.weight.data.device) + + def eject_loras(self, motion_module: MotionWrapper, lora_stack: List[Tuple[float, Dict[str, Tensor]]]): + lora_stack.reverse() # should not matter but just in case + for lora in lora_stack: + (state_dict, alpha) = lora + + for key in state_dict: + layer_infos = key.split(".") + + curr_layer = motion_module + while len(layer_infos) > 0: + temp_name = layer_infos.pop(0) + curr_layer = curr_layer.__getattr__(temp_name) + + curr_layer.weight.data -= alpha * state_dict[key].to(curr_layer.weight.data.device) + def load_motion_module( self, model_name: str, + lora_stack: List = None, ): motion_module = load_motion_module(model_name) + # inject loras + if motion_module.is_v2: + if hasattr(motion_module, "lora_stack") and isinstance(motion_module.lora_stack, list): + self.eject_loras(motion_module, motion_module.lora_stack) + delattr(motion_module, "lora_stack") + + if isinstance(lora_stack, list): + self.inject_loras(motion_module, lora_stack) + setattr(motion_module, "lora_stack", lora_stack) + + elif isinstance(lora_stack, list): + logger.warning("LoRA is provided but only motion module v2 is supported.") + return (motion_module,) +class AnimateDiffLoraLoader: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "lora_name": (get_available_loras(),), + "alpha": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + }, + "optional": { + "lora_stack": ("MOTION_LORA_STACK",), + }, + } + + RETURN_TYPES = ("MOTION_LORA_STACK",) + CATEGORY = "Animate Diff" + FUNCTION = "load_lora" + + def load_lora( + self, + lora_name: str, + alpha: float, + lora_stack: List = None, + ): + if not lora_stack: + lora_stack = [] + + lora = load_lora(lora_name) + lora_stack.append((lora, alpha)) + + return (lora_stack,) + + class AnimateDiffCombine: @classmethod def INPUT_TYPES(s): @@ -324,6 +404,7 @@ class ImageChunking: NODE_CLASS_MAPPINGS = { "AnimateDiffModuleLoader": AnimateDiffModuleLoader, + "AnimateDiffLoraLoader": AnimateDiffLoraLoader, "AnimateDiffCombine": AnimateDiffCombine, "AnimateDiffSampler": AnimateDiffSampler, "AnimateDiffSlidingWindowOptions": AnimateDiffSlidingWindowOptions, @@ -332,6 +413,7 @@ NODE_CLASS_MAPPINGS = { } NODE_DISPLAY_NAME_MAPPINGS = { "AnimateDiffModuleLoader": "Animate Diff Module Loader", + "AnimateDiffLoraLoader": "Animate Diff Lora Loader", "AnimateDiffSampler": "Animate Diff Sampler", "AnimateDiffSlidingWindowOptions": "Sliding Window Options", "AnimateDiffCombine": "Animate Diff Combine", diff --git a/loras/.gitkeep b/loras/.gitkeep new file mode 100644 index 0000000..e69de29 diff --git a/workflows/lora.json b/workflows/lora.json new file mode 100644 index 0000000..72efff8 --- /dev/null +++ b/workflows/lora.json @@ -0,0 +1,515 @@ +{ + "last_node_id": 21, + "last_link_id": 38, + "nodes": [ + { + "id": 6, + "type": "CLIPTextEncode", + "pos": [ + 415, + 186 + ], + "size": { + "0": 422.84503173828125, + "1": 164.31304931640625 + }, + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 3 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 29 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "photo of coastline, rocks, storm weather, wind, waves, lightning, 8k uhd, dslr, soft lighting, high quality, film grain, Fujifilm XT3" + ] + }, + { + "id": 8, + "type": "VAEDecode", + "pos": [ + 1253, + 191 + ], + "size": { + "0": 210, + "1": 46 + }, + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "samples", + "type": "LATENT", + "link": 28 + }, + { + "name": "vae", + "type": "VAE", + "link": 20 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 19 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "VAEDecode" + } + }, + { + "id": 12, + "type": "AnimateDiffCombine", + "pos": [ + 1254, + 290 + ], + "size": [ + 315, + 507 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 19 + } + ], + "properties": { + "Node name for S&R": "AnimateDiffCombine" + }, + "widgets_values": [ + 8, + 0, + false, + "AnimateDiff", + "image/gif", + false, + "/view?filename=AnimateDiff_00003_.gif&subfolder=&type=temp&format=image%2Fgif" + ] + }, + { + "id": 7, + "type": "CLIPTextEncode", + "pos": [ + 413, + 389 + ], + "size": { + "0": 425.27801513671875, + "1": 180.6060791015625 + }, + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 5 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 30 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "blur, haze, deformed iris, deformed pupils, semi-realistic, cgi, 3d, render, sketch, cartoon, drawing, anime, mutated hands and fingers, deformed, distorted, disfigured, poorly drawn, bad anatomy, wrong anatomy, extra limb, missing limb, floating limbs, disconnected limbs, mutation, mutated, ugly, disgusting, amputation" + ] + }, + { + "id": 20, + "type": "EmptyLatentImage", + "pos": [ + 522, + 621 + ], + "size": { + "0": 315, + "1": 106 + }, + "flags": {}, + "order": 0, + "mode": 0, + "outputs": [ + { + "name": "LATENT", + "type": "LATENT", + "links": [ + 35 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "EmptyLatentImage" + }, + "widgets_values": [ + 512, + 512, + 1 + ] + }, + { + "id": 13, + "type": "VAELoader", + "pos": [ + 28, + 223 + ], + "size": { + "0": 315, + "1": 58 + }, + "flags": {}, + "order": 1, + "mode": 0, + "outputs": [ + { + "name": "VAE", + "type": "VAE", + "links": [ + 20 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "VAELoader" + }, + "widgets_values": [ + "vae-ft-mse-840000-ema-pruned.safetensors" + ] + }, + { + "id": 15, + "type": "AnimateDiffSampler", + "pos": [ + 882, + 192 + ], + "size": { + "0": 315, + "1": 350 + }, + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [ + { + "name": "motion_module", + "type": "MOTION_MODULE", + "link": 24, + "slot_index": 0 + }, + { + "name": "model", + "type": "MODEL", + "link": 25, + "slot_index": 1 + }, + { + "name": "positive", + "type": "CONDITIONING", + "link": 29 + }, + { + "name": "negative", + "type": "CONDITIONING", + "link": 30 + }, + { + "name": "latent_image", + "type": "LATENT", + "link": 35 + }, + { + "name": "sliding_window_opts", + "type": "SLIDING_WINDOW_OPTS", + "link": null + } + ], + "outputs": [ + { + "name": "LATENT", + "type": "LATENT", + "links": [ + 28 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "AnimateDiffSampler" + }, + "widgets_values": [ + "default", + 14, + 45987230, + "fixed", + 25, + 7.5, + "ddim", + "ddim_uniform", + 1 + ] + }, + { + "id": 16, + "type": "AnimateDiffModuleLoader", + "pos": [ + 27, + 345 + ], + "size": { + "0": 315, + "1": 58 + }, + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "lora_stack", + "type": "MOTION_LORA_STACK", + "link": 38, + "slot_index": 0 + } + ], + "outputs": [ + { + "name": "MOTION_MODULE", + "type": "MOTION_MODULE", + "links": [ + 24 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "AnimateDiffModuleLoader" + }, + "widgets_values": [ + "mm_sd_v15_v2.ckpt" + ] + }, + { + "id": 21, + "type": "AnimateDiffLoraLoader", + "pos": [ + -317, + 350 + ], + "size": [ + 310, + 80 + ], + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [ + { + "name": "lora_stack", + "type": "MOTION_LORA_STACK", + "link": null + } + ], + "outputs": [ + { + "name": "MOTION_LORA_STACK", + "type": "MOTION_LORA_STACK", + "links": [ + 38 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "AnimateDiffLoraLoader" + }, + "widgets_values": [ + "v2_lora_ZoomIn.ckpt", + 1 + ] + }, + { + "id": 4, + "type": "CheckpointLoaderSimple", + "pos": [ + 28, + 457 + ], + "size": { + "0": 315, + "1": 98 + }, + "flags": {}, + "order": 2, + "mode": 0, + "outputs": [ + { + "name": "MODEL", + "type": "MODEL", + "links": [ + 25 + ], + "slot_index": 0 + }, + { + "name": "CLIP", + "type": "CLIP", + "links": [ + 3, + 5 + ], + "slot_index": 1 + }, + { + "name": "VAE", + "type": "VAE", + "links": [], + "slot_index": 2 + } + ], + "properties": { + "Node name for S&R": "CheckpointLoaderSimple" + }, + "widgets_values": [ + "RealisticVision_v20.safetensors" + ] + } + ], + "links": [ + [ + 3, + 4, + 1, + 6, + 0, + "CLIP" + ], + [ + 5, + 4, + 1, + 7, + 0, + "CLIP" + ], + [ + 19, + 8, + 0, + 12, + 0, + "IMAGE" + ], + [ + 20, + 13, + 0, + 8, + 1, + "VAE" + ], + [ + 24, + 16, + 0, + 15, + 0, + "MOTION_MODULE" + ], + [ + 25, + 4, + 0, + 15, + 1, + "MODEL" + ], + [ + 28, + 15, + 0, + 8, + 0, + "LATENT" + ], + [ + 29, + 6, + 0, + 15, + 2, + "CONDITIONING" + ], + [ + 30, + 7, + 0, + 15, + 3, + "CONDITIONING" + ], + [ + 35, + 20, + 0, + 15, + 4, + "LATENT" + ], + [ + 38, + 21, + 0, + 16, + 0, + "MOTION_LORA_STACK" + ] + ], + "groups": [], + "config": {}, + "extra": {}, + "version": 0.4 +} \ No newline at end of file