initial lora support
This commit is contained in:
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user