Added Motion LoRA support, updated Loader node to support AnimateDiff LoRA Loader input

This commit is contained in:
Jedrzej Kosinski
2023-09-25 10:28:38 -05:00
parent 53159c3067
commit 4243d40a88
6 changed files with 206 additions and 16 deletions
+20 -2
View File
@@ -71,8 +71,10 @@ class BetaScheduleCache:
class Folders:
ANIMATEDIFF_MODELS = "AnimateDiffEvolved_Models"
MOTION_LORA = "AnimateDiffMotion_LoRA"
# register motion models folder(s)
folder_paths.folder_names_and_paths[Folders.ANIMATEDIFF_MODELS] = (
[
str(Path(__file__).parent.parent / "models")
@@ -80,6 +82,14 @@ folder_paths.folder_names_and_paths[Folders.ANIMATEDIFF_MODELS] = (
folder_paths.supported_pt_extensions
)
# register motion LoRA folder(s)
folder_paths.folder_names_and_paths[Folders.MOTION_LORA] = (
[
str(Path(__file__).parent.parent / "motion_lora")
],
folder_paths.supported_pt_extensions
)
#Register video_formats folder
folder_paths.folder_names_and_paths["video_formats"] = (
@@ -98,8 +108,16 @@ def get_motion_model_path(model_name: str):
return folder_paths.get_full_path(Folders.ANIMATEDIFF_MODELS, model_name)
def get_available_motion_loras():
return folder_paths.get_filename_list(Folders.MOTION_LORA)
def get_motion_lora_path(lora_name: str):
return folder_paths.get_full_path(Folders.MOTION_LORA, lora_name)
# modified from https://stackoverflow.com/questions/22058048/hashing-a-file-in-python
def calculate_file_hash(filename: str):
def calculate_file_hash(filename: str, hash_every_n: int = 50):
h = hashlib.sha256()
b = bytearray(1024*1024)
mv = memoryview(b)
@@ -107,7 +125,7 @@ def calculate_file_hash(filename: str):
i = 0
# don't hash entire file, only portions of it
while n := f.readinto(mv):
if i%50 == 0:
if i%hash_every_n == 0:
h.update(mv[:n])
i += 1
return h.hexdigest()
+39
View File
@@ -0,0 +1,39 @@
from torch import Tensor
from pathlib import Path
class MotionLoRAInfo:
def __init__(self, name: str, strength: float = 1.0, hash: str=""):
self.name = name
self.strength = strength
self.hash = ""
def set_hash(self, hash: str):
self.hash = hash
def clone(self):
return MotionLoRAInfo(self.name, self.strength, self.hash)
class MotionLoRAWrapper:
def __init__(self, state_dict: dict[str, Tensor], hash: str):
self.state_dict = state_dict
self.hash = hash
self.info: MotionLoRAInfo = None
def set_info(self, info: MotionLoRAInfo):
self.info = info
class MotionLoRAList:
def __init__(self):
self.loras: list[MotionLoRAInfo] = []
def add_lora(self, lora: MotionLoRAInfo):
self.loras.append(lora)
def clone(self):
new_list = MotionLoRAList()
for lora in self.loras:
new_list.add_lora(lora.clone())
return new_list
+99 -6
View File
@@ -12,7 +12,8 @@ import comfy.model_management as model_management
from comfy.cli_args import args
from comfy.utils import calculate_parameters, load_torch_file
from .model_utils import calculate_file_hash, get_motion_model_path, is_checkpoint_sd1_5
from .motion_lora import MotionLoRAList, MotionLoRAWrapper, MotionLoRAInfo
from .model_utils import calculate_file_hash, get_motion_lora_path, get_motion_model_path, is_checkpoint_sd1_5
from .logger import logger
@@ -41,21 +42,97 @@ def clone_injection(self, *args, **kwargs):
comfy_model_patcher.ModelPatcher.clone = clone_injection
# cache loaded motion modules
# cached motion modules
motion_modules: dict[str, 'MotionWrapper'] = {}
# cached motion loras
motion_loras: dict[str, MotionLoRAWrapper] = {}
def load_motion_module(model_name: str):
# adapted from https://github.com/guoyww/AnimateDiff/blob/main/animatediff/utils/convert_lora_safetensor_to_diffusers.py
# Example LoRA keys:
# down_blocks.0.motion_modules.0.temporal_transformer.transformer_blocks.0.attention_blocks.0.processor.to_q_lora.down.weight
# down_blocks.0.motion_modules.0.temporal_transformer.transformer_blocks.0.attention_blocks.0.processor.to_q_lora.up.weight
#
# Example model keys:
# down_blocks.0.motion_modules.0.temporal_transformer.transformer_blocks.0.attention_blocks.0.to_q.weight
#
def apply_lora_to_mm_state_dict(model_dict: dict[str, Tensor], lora: MotionLoRAWrapper):
model_has_midblock = has_mid_block(model_dict)
lora_has_midblock = has_mid_block(lora.state_dict)
def get_version(has_midblock: bool):
return "v2" if has_midblock else "v1"
logger.info(f"Applying a {get_version(lora_has_midblock)} LoRA ({lora.info.name}) to a {get_version(model_has_midblock)} motion model.")
for key in lora.state_dict:
# if motion model doesn't have a mid_block, skip mid_block entries
if not model_has_midblock:
if "mid_block" in key: continue
# only process lora down key (we will process up at the same time as down)
if "up." in key: continue
# key to get up value
up_key = key.replace(".down.", ".up.")
# adapt key to match model_dict format - remove 'processor.', '_lora', 'down.', and 'up.'
model_key = key.replace("processor.", "").replace("_lora", "").replace("down.", "").replace("up.", "")
# model keys have a '0.' after all 'to_out.' weight keys
model_key = model_key.replace("to_out.", "to_out.0.")
weight_down = lora.state_dict[key]
weight_up = lora.state_dict[up_key]
# apply weights to model_dict - multiply strength by matrix multiplication of up and down weights
model_dict[model_key] += lora.info.strength * torch.mm(weight_up, weight_down).to(model_dict[model_key].device)
def load_motion_lora(lora_name: str) -> MotionLoRAWrapper:
# if already loaded, return it
lora_path = get_motion_lora_path(lora_name)
lora_hash = calculate_file_hash(lora_path, hash_every_n=3)
if lora_hash in motion_loras:
return motion_loras[lora_hash]
logger.info(f"Loading motion LoRA {lora_name}")
l_state_dict = load_torch_file(lora_path)
lora = MotionLoRAWrapper(l_state_dict, lora_hash)
# add motion LoRA to cache
motion_loras[lora_hash] = lora
return lora
def load_motion_module(model_name: str, motion_lora: MotionLoRAList = None) -> 'MotionWrapper':
# if already loaded, return it
model_path = get_motion_model_path(model_name)
model_hash = calculate_file_hash(model_path)
model_hash = calculate_file_hash(model_path, hash_every_n=50)
# load lora, if present
loras = []
if motion_lora is not None:
for lora_info in motion_lora.loras:
lora = load_motion_lora(lora_info.name)
lora.set_info(lora_info)
loras.append(lora)
loras.sort(key=lambda x: x.hash)
# use lora hashes with model hash
for lora in loras:
model_hash += lora.hash
model_hash = str(hash(model_hash))
# models are determined by combo self + applied loras
if model_hash in motion_modules:
return motion_modules[model_hash]
logger.info(f"Loading motion module {model_name}")
mm_state_dict = load_torch_file(model_path)
motion_module = MotionWrapper(mm_state_dict=mm_state_dict, mm_name=model_name)
# load lora state dicts if exist
if len(loras) > 0:
for lora in loras:
# apply LoRA to mm_state_dict
apply_lora_to_mm_state_dict(mm_state_dict, lora)
motion_module = MotionWrapper(mm_state_dict=mm_state_dict, mm_hash=model_hash, mm_name=model_name, loras=loras)
parameters = calculate_parameters(mm_state_dict, "")
usefp16 = model_management.should_use_fp16(model_params=parameters)
@@ -71,6 +148,11 @@ def load_motion_module(model_name: str):
return motion_module
def unload_motion_module(motion_module: 'MotionWrapper'):
logger.info(f"Removing motion module {motion_module.mm_name} from cache")
motion_modules.pop(motion_module.mm_hash, None)
##################################################################################
##################################################################################
# Injection-related classes and functions
@@ -203,6 +285,7 @@ class InjectionParams:
self.context_schedule: str = None
self.closed_loop: bool = False
self.version: str = None
self.loras: MotionLoRAList = None
def set_version(self, motion_module: 'MotionWrapper'):
self.version = motion_module.version
@@ -214,6 +297,9 @@ class InjectionParams:
self.context_schedule = context_schedule
self.closed_loop = closed_loop
def set_loras(self, loras: MotionLoRAList):
self.loras = loras.clone()
def reset_context(self):
self.context_length = None
self.context_stride = None
@@ -232,6 +318,8 @@ class InjectionParams:
context_overlap=self.context_overlap, context_schedule=self.context_schedule,
closed_loop=self.closed_loop
)
if self.loras is not None:
new_params.loras = self.loras.clone()
return new_params
@@ -305,7 +393,7 @@ def has_mid_block(mm_state_dict: dict[str, Tensor]):
class MotionWrapper(nn.Module):
def __init__(self, mm_state_dict: dict[str, Tensor], mm_name: str="mm_sd_v15.ckpt"):
def __init__(self, mm_state_dict: dict[str, Tensor], mm_hash: str, mm_name: str="mm_sd_v15.ckpt" , loras: list[MotionLoRAInfo]=None):
super().__init__()
self.down_blocks = nn.ModuleList([])
self.up_blocks = nn.ModuleList([])
@@ -317,9 +405,14 @@ class MotionWrapper(nn.Module):
self.up_blocks.append(MotionModule(c, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.UP))
if has_mid_block(mm_state_dict):
self.mid_block = MotionModule(1280, temporal_position_encoding_max_len=self.encoding_max_len, block_type=BlockType.MID)
self.mm_hash = mm_hash
self.mm_name = mm_name
self.version = "v1" if self.mid_block is None else "v2"
self.AD_video_length: int = 24
self.loras = loras
def has_loras(self):
return self.loras is not None
def set_video_length(self, video_length: int):
self.AD_video_length = video_length
+42 -6
View File
@@ -14,9 +14,10 @@ import folder_paths
from comfy.sd import load_checkpoint_guess_config
from comfy.model_patcher import ModelPatcher
from .logger import logger
from .motion_module import InjectorVersion, eject_params_from_model, inject_params_into_model, load_motion_module
from .motion_module import InjectionParams
from .model_utils import IsChangedHelper, get_available_motion_models, BetaSchedules, raise_if_not_checkpoint_sd1_5
from .motion_lora import MotionLoRAInfo, MotionLoRAList
from .motion_module import InjectorVersion, eject_params_from_model, get_injected_mm_params, inject_params_into_model, load_motion_lora, load_motion_module
from .motion_module import InjectionParams, is_injected_mm_params
from .model_utils import IsChangedHelper, get_available_motion_loras, get_available_motion_models, BetaSchedules, raise_if_not_checkpoint_sd1_5
from .context import ContextOptions, ContextSchedules, UniformContextOptions
from .sampling import animatediff_sample_factory
@@ -33,6 +34,36 @@ sys.path.insert(0, Path(__file__).parent.parent.parent.parent)
import nodes as comfy_nodes
class AnimateDiffLoraLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"lora_name": (get_available_motion_loras(),),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0}),
},
"optional": {
"prev_motion_lora": ("MOTION_LORA",),
}
}
RETURN_TYPES = ("MOTION_LORA",)
CATEGORY = "Animate Diff"
FUNCTION = "load_motion_lora"
def load_motion_lora(self, lora_name: str, strength: float, prev_motion_lora: MotionLoRAList=None):
if prev_motion_lora is None:
prev_motion_lora = MotionLoRAList()
else:
prev_motion_lora = prev_motion_lora.clone()
# load lora
lora = load_motion_lora(lora_name)
lora_info = MotionLoRAInfo(name=lora_name, strength=strength, hash=lora.hash)
prev_motion_lora.add_lora(lora_info)
return (prev_motion_lora,)
class AnimateDiffLoaderWithContext:
@classmethod
def INPUT_TYPES(s):
@@ -43,7 +74,8 @@ class AnimateDiffLoaderWithContext:
"beta_schedule": (BetaSchedules.get_alias_list_with_first_element(BetaSchedules.SQRT_LINEAR),),
},
"optional": {
"context_options": ("CONTEXT_OPTIONS",)
"context_options": ("CONTEXT_OPTIONS",),
"motion_lora": ("MOTION_LORA",),
}
}
@@ -55,11 +87,11 @@ class AnimateDiffLoaderWithContext:
def load_mm_and_inject_params(self,
model: ModelPatcher,
model_name: str, beta_schedule: str,
context_options: ContextOptions=None,
context_options: ContextOptions=None, motion_lora: MotionLoRAList=None,
):
raise_if_not_checkpoint_sd1_5(model)
# load motion module
load_motion_module(model_name)
load_motion_module(model_name, motion_lora)
# set injection params
injection_params = InjectionParams(
video_length=None,
@@ -78,6 +110,8 @@ class AnimateDiffLoaderWithContext:
context_schedule=context_options.context_schedule,
closed_loop=context_options.closed_loop
)
if motion_lora:
injection_params.set_loras(motion_lora)
# inject for use in sampling code
model = inject_params_into_model(model, injection_params)
@@ -408,6 +442,7 @@ class EmptyLatentImageLarge:
NODE_CLASS_MAPPINGS = {
"ADE_AnimateDiffUniformContextOptions": AnimateDiffUniformContextOptions,
"ADE_AnimateDiffLoaderWithContext": AnimateDiffLoaderWithContext,
"ADE_AnimateDiffLoRALoader": AnimateDiffLoraLoader,
"ADE_AnimateDiffUnload": AnimateDiffUnload,
"ADE_AnimateDiffCombine": AnimateDiffCombine,
"ADE_EmptyLatentImageLarge": EmptyLatentImageLarge,
@@ -418,6 +453,7 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"ADE_AnimateDiffUniformContextOptions": "Uniform Context Options",
"ADE_AnimateDiffLoaderWithContext": "AnimateDiff Loader",
"ADE_AnimateDiffLoRALoader": "AnimateDiff LoRA Loader",
"ADE_AnimateDiffUnload": "AnimateDiff Unload",
"ADE_AnimateDiffCombine": "AnimateDiff Combine",
"ADE_EmptyLatentImageLarge": "Empty Latent Image (Big Batch)",
+6 -2
View File
@@ -18,7 +18,7 @@ import comfy.ldm.modules.diffusionmodules.openaimodel as openaimodel
import comfy.model_management as model_management
from .logger import logger
from .motion_module import InjectionParams, VanillaTemporalModule, eject_motion_module, inject_motion_module, inject_params_into_model, load_motion_module
from .motion_module import InjectionParams, VanillaTemporalModule, eject_motion_module, inject_motion_module, inject_params_into_model, load_motion_module, unload_motion_module
from .motion_module import is_injected_mm_params, get_injected_mm_params
from .context import get_context_scheduler
from .model_utils import BetaScheduleCache, BetaSchedules, wrap_function_to_inject_xformers_bug_info
@@ -136,7 +136,7 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable:
##############################################
# try to load motion module
motion_module = load_motion_module(params.model_name)
motion_module = load_motion_module(params.model_name, params.loras)
# inject motion module into unet
inject_motion_module(model=model, motion_module=motion_module, params=params)
@@ -162,6 +162,10 @@ def animatediff_sample_factory(orig_comfy_sample: Callable) -> Callable:
finally:
# attempt to eject motion module
eject_motion_module(model=model)
# if loras are present, remove model so it can be re-loaded next time with fresh weights
if motion_module.has_loras():
unload_motion_module(motion_module)
del motion_module
##############################################
# Restoration
model_management.maximum_batch_area = orig_maximum_batch_area
View File