Readded motion LoRA support
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
@@ -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] = []
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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,)
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user