Readded motion LoRA support

This commit is contained in:
Jedrzej Kosinski
2023-12-06 07:49:19 -06:00
parent 59566d1eae
commit 63e5ee26ae
6 changed files with 28 additions and 35 deletions
+7 -5
View File
@@ -14,9 +14,9 @@ from .motion_module_ad import AnimateDiffModelWrapper, has_mid_block, normalize_
from .logger import logger
from .motion_utils import MotionCompatibilityError, NoiseType, normalize_min_max
from .motion_lora import MotionLoraInfo, MotionLoraList, MotionLoraWrapper
from .motion_lora import MotionLoraInfo, MotionLoraList
from .model_utils import ModelTypesSD, calculate_file_hash, get_motion_lora_path, get_motion_model_path, \
from .model_utils import ModelTypeSD, calculate_file_hash, get_motion_lora_path, get_motion_model_path, \
get_sd_model_type
@@ -55,15 +55,15 @@ class ModelPatcherAndInjector(ModelPatcher):
return super().unpatch_model(device_to)
def inject_model(self, device_to=None):
# motion_model.model is AnimateDiffModelWrapper
if self.motion_model is not None:
self.motion_model.model.eject(self)
self.motion_model.model.inject(self)
self.motion_model.model.to(device_to)
def eject_model(self, device_to=None):
# motion_model.model is AnimateDiffModelWrapper
if self.motion_model is not None:
self.motion_model.model.eject(self)
self.motion_model.model.to(device_to)
def clone(self):
cloned = ModelPatcherAndInjector(self)
@@ -73,6 +73,7 @@ class ModelPatcherAndInjector(ModelPatcher):
class MotionModelPatcher(ModelPatcher):
# Mostly here so that type hints work in IDEs
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.model: AnimateDiffModelWrapper = self.model
@@ -118,7 +119,8 @@ def load_motion_lora_as_patches(motion_model: MotionModelPatcher, lora: MotionLo
weight_down = state_dict[key]
weight_up = state_dict[up_key]
# actual weights obtained by matrix multiplication of up and down weights
patches[model_key] += torch.mm(weight_up, weight_down)
# save as a tuple, so that (Motion)ModelPatcher's calculate_weight function detects len==1, applying it correctly
patches[model_key] = (torch.mm(weight_up, weight_down),)
del state_dict
# add patches to motion ModelPatcher
motion_model.add_patches(patches=patches, strength_patch=lora.strength)
+6 -6
View File
@@ -154,7 +154,7 @@ def calculate_model_hash(model: ModelPatcher):
return m.hexdigest()
class ModelTypesSD:
class ModelTypeSD:
SD1_5 = "SD1.5"
SD2_1 = "SD2.1"
SDXL = "SDXL"
@@ -166,15 +166,15 @@ def get_sd_model_type(model: ModelPatcher) -> str:
if model is None:
return None
elif type(model.model) == BaseModel:
return ModelTypesSD.SD1_5
return ModelTypeSD.SD1_5
elif type(model.model) == SDXL:
return ModelTypesSD.SDXL
return ModelTypeSD.SDXL
elif type(model.model) == SD21UNCLIP:
return ModelTypesSD.SD2_1
return ModelTypeSD.SD2_1
elif type(model.model) == SDXLRefiner:
return ModelTypesSD.SDXL_REFINER
return ModelTypeSD.SDXL_REFINER
elif type(model.model) == SVD_img2vid:
return ModelTypesSD.SVD
return ModelTypeSD.SVD
else:
return str(type(model.model).__name__)
-13
View File
@@ -1,6 +1,3 @@
from torch import Tensor
class MotionLoraInfo:
def __init__(self, name: str, strength: float = 1.0, hash: str=""):
self.name = name
@@ -14,16 +11,6 @@ class MotionLoraInfo:
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] = []
+4 -4
View File
@@ -12,7 +12,7 @@ from ldm.modules.diffusionmodules import openaimodel
from comfy.ldm.modules.diffusionmodules.openaimodel import ResBlock, SpatialTransformer
from .motion_lora import MotionLoraInfo
from .motion_utils import GenericMotionWrapper, GroupNormAD, InjectorVersion, BlockType, CrossAttentionMM, MotionCompatibilityError, TemporalTransformerGeneric, prepare_mask_batch
from .model_utils import ModelTypesSD
from .model_utils import ModelTypeSD
def zero_module(module):
@@ -96,9 +96,9 @@ def normalize_ad_state_dict(mm_state_dict: dict[str, Tensor], mm_name: str) -> T
sd_type: str = None
down_block_max = get_down_block_max(mm_state_dict)
if down_block_max == 3:
sd_type = ModelTypesSD.SD1_5
sd_type = ModelTypeSD.SD1_5
elif down_block_max == 2:
sd_type = ModelTypesSD.SDXL
sd_type = ModelTypeSD.SDXL
else:
raise ValueError(f"'{mm_name}' is not a valid SD1.5 nor SDXL motion module - contained {down_block_max} downblocks.")
# determine the model's format
@@ -141,7 +141,7 @@ class AnimateDiffModelWrapper(nn.Module):
self.mid_block: Union[MotionModule, None] = None
self.encoding_max_len = get_position_encoding_max_len(mm_state_dict, mm_info)
# SDXL has 3 up/down blocks, SD1.5 has 4 up/down blocks
if mm_info.sd_type == ModelTypesSD.SDXL:
if mm_info.sd_type == ModelTypeSD.SDXL:
layer_channels = (320, 640, 1280)
else:
layer_channels = (320, 640, 1280, 1280)
+9 -5
View File
@@ -1,3 +1,4 @@
from pathlib import Path
import torch
import comfy.sample as comfy_sample
@@ -5,7 +6,7 @@ from comfy.model_patcher import ModelPatcher
from .context import ContextOptions, ContextSchedules, UniformContextOptions
from .logger import logger
from .model_utils import get_available_motion_loras, get_available_motion_models, BetaSchedules
from .model_utils import get_available_motion_loras, get_available_motion_models, BetaSchedules, get_motion_lora_path
from .motion_utils import NoiseType
from .motion_lora import MotionLoraInfo, MotionLoraList
from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelSettings, load_motion_module
@@ -69,10 +70,13 @@ class AnimateDiffLoraLoader:
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)
# check if motion lora with name exists
lora_path = get_motion_lora_path(lora_name)
if not Path(lora_path).is_file():
raise FileNotFoundError(f"Motion lora with name '{lora_name}' not found.")
# create motion lora info to be loaded in AnimateDiff Loader
lora_info = MotionLoraInfo(name=lora_name, strength=strength)
prev_motion_lora.add_lora(lora_info)
return (prev_motion_lora,)
+2 -2
View File
@@ -16,7 +16,7 @@ from comfy.controlnet import ControlBase
from .context import get_context_scheduler
from .motion_utils import GroupNormAD, NoiseType
from .model_utils import BetaScheduleCache, BetaSchedules, ModelTypesSD, wrap_function_to_inject_xformers_bug_info
from .model_utils import BetaScheduleCache, BetaSchedules, ModelTypeSD, wrap_function_to_inject_xformers_bug_info
from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelPatcher
from .motion_module_ad import AnimateDiffFormat, AnimateDiffInfo, AnimateDiffModelWrapper, VanillaTemporalModule
from .logger import logger
@@ -219,7 +219,7 @@ def motion_sample_factory(orig_comfy_sample: Callable) -> Callable:
model.model.memory_required = unlimited_memory_required
# only apply groupnorm hack if not [AnimateDiff SD1.5 and v2 and should apply v2 properly]
info: AnimateDiffInfo = model.motion_model.model.mm_info
if not (info.mm_format == AnimateDiffFormat.ANIMATEDIFF and info.sd_type == ModelTypesSD.SD1_5 and \
if not (info.mm_format == AnimateDiffFormat.ANIMATEDIFF and info.sd_type == ModelTypeSD.SD1_5 and \
info.mm_version == "v2" and params.apply_v2_models_properly):
torch.nn.GroupNorm.forward = groupnorm_mm_factory(params)
if params.apply_mm_groupnorm_hack: