initial lora support

This commit is contained in:
Tung Nguyen
2023-09-25 21:26:51 +07:00
parent a0bdb7e06c
commit d5d42f8f3e
4 changed files with 127 additions and 3 deletions
+44
View File
@@ -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,16 +22,31 @@ 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
@@ -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)
motion_loras[lora_hash] = updated_state_dict
return motion_loras[lora_hash]
+36 -2
View File
@@ -3,14 +3,14 @@ import json
import torch
import numpy as np
import hashlib
from typing import List
from typing import List, Dict
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 .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
@@ -43,6 +43,38 @@ class AnimateDiffModuleLoader:
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 +356,7 @@ class ImageChunking:
NODE_CLASS_MAPPINGS = {
"AnimateDiffModuleLoader": AnimateDiffModuleLoader,
"AnimateDiffLoraLoader": AnimateDiffLoraLoader,
"AnimateDiffCombine": AnimateDiffCombine,
"AnimateDiffSampler": AnimateDiffSampler,
"AnimateDiffSlidingWindowOptions": AnimateDiffSlidingWindowOptions,
@@ -332,6 +365,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",
+47 -1
View File
@@ -2,6 +2,7 @@ import torch
from torch import Tensor
from torch.nn.functional import group_norm
from einops import rearrange
from typing import List, Tuple, Dict
import comfy.ldm.modules.diffusionmodules.openaimodel as openaimodel
import comfy.model_management as model_management
@@ -173,7 +174,10 @@ class AnimateDiffSampler(KSampler):
}
}
inputs["required"].update(KSampler.INPUT_TYPES()["required"])
inputs["optional"] = {"sliding_window_opts": ("SLIDING_WINDOW_OPTS",)}
inputs["optional"] = {
"sliding_window_opts": ("SLIDING_WINDOW_OPTS",),
"lora_stack": ("MOTION_LORA_STACK",),
}
return inputs
FUNCTION = "animatediff_sample"
@@ -229,6 +233,20 @@ class AnimateDiffSampler(KSampler):
inject_sampling_function(ctx)
def inject_loras(self, model, lora_stack: List[Tuple[float, Dict[str, Tensor]]]):
for lora in lora_stack:
(alpha, state_dict) = lora
for key in state_dict:
layer_infos = key.split(".")
curr_layer = model.diffusion_model
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_motion_module(self, model, inject_method):
unet = model.model.diffusion_model
@@ -244,6 +262,20 @@ class AnimateDiffSampler(KSampler):
def eject_sliding_sampler(self):
eject_sampling_function()
def eject_loras(self, model, lora_stack: List[Tuple[float, Dict[str, Tensor]]]):
for lora in lora_stack.reverse():
(alpha, state_dict) = lora
for key in state_dict:
layer_infos = key.split(".")
curr_layer = model.diffusion_model
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 animatediff_sample(
self,
motion_module,
@@ -260,6 +292,7 @@ class AnimateDiffSampler(KSampler):
latent_image,
denoise=1.0,
sliding_window_opts: SlidingContext = None,
lora_stack: List = None,
**kwargs,
):
# init latents
@@ -287,6 +320,15 @@ class AnimateDiffSampler(KSampler):
# inject motion module
model = self.inject_motion_module(model, motion_module, inject_method, video_length)
# inject loras
if isinstance(lora_stack, list):
if motion_module.is_v2:
self.inject_loras(model, lora_stack)
else:
logger.warning(
"Lora is provided but only motion module v2 is supported. Please switch to mm_sd_v15_v2 module."
)
# inject sliding sampler
if is_sliding:
self.inject_sliding_sampler(frame_number, sliding_window_opts=sliding_window_opts)
@@ -308,6 +350,10 @@ class AnimateDiffSampler(KSampler):
except:
raise
finally:
# eject loras
if isinstance(lora_stack, list) and motion_module.is_v2:
self.eject_loras(model, lora_stack)
# eject motion module
self.eject_motion_module(model, inject_method)
View File