Added Motion LoRA support, updated Loader node to support AnimateDiff LoRA Loader input
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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)",
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user