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`
+
+
+
+Example workflow:
+
+
+Workflow: [lora.json](https://github.com/ArtVentureX/comfyui-animatediff/blob/main/workflows/lora.json)
+
+Samples:
+
+
+
+
+
+ |
+
+
+
+
+ |
+
+
+
+
+ |
+
+
+
+
+ |
+
+
+
## 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