From 2ae255e97d29379518d293eca80d647609f3c14e Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 23 Mar 2024 12:12:01 -0500 Subject: [PATCH] Added lora masking support for CLIP component of lora --- animatediff/model_injection.py | 38 ++++++++++++++++++++++-- animatediff/nodes.py | 16 +++++------ animatediff/nodes_conditioning.py | 25 ++++++++++++++-- animatediff/nodes_lora.py | 48 ------------------------------- 4 files changed, 66 insertions(+), 61 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 371eaf5..da0fe6d 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -255,6 +255,39 @@ class ModelPatcherAndInjector(ModelPatcher): return ModelPatcherAndInjector(model) +class CLIPWithHooks(CLIP): + def __init__(self, clip: Union[CLIP, 'CLIPWithHooks']): + super().__init__(no_init=True) + self.patcher = ModelPatcherAndInjector.create_from(clip.patcher) + self.cond_stage_model = clip.cond_stage_model + self.tokenizer = clip.tokenizer + self.layer_idx = clip.layer_idx + self.desired_hooks: LoraHookGroup = None + if hasattr(clip, "desired_hooks"): + self.desired_hooks = clip.desired_hooks + + def clone(self): + cloned = CLIPWithHooks(clip=self) + return cloned + + def set_desired_hooks(self, lora_hooks: LoraHookGroup): + self.desired_hooks = lora_hooks + + def add_hooked_patches(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): + return self.patcher.add_hooked_patches(lora_hook=lora_hook, patches=patches, strength_patch=strength_patch, strength_model=strength_model) + + def load_model(self, *args, **kwargs): + self.patcher.unpatch_hooked() + returned = super().load_model(*args, **kwargs) + # apply desired hooks + self.patcher.patch_hooked(lora_hooks=self.desired_hooks) + return returned + + def encode_from_tokens(self, *args, **kwargs): + returned = super().encode_from_tokens(*args, **kwargs) + return returned + + def load_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInjector], clip: CLIP, lora: dict[str, Tensor], lora_hook: LoraHook, strength_model: float, strength_clip: float): key_map = {} if model is not None: @@ -271,9 +304,8 @@ def load_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInject new_modelpatcher = None if clip is not None: - new_clip = clip.clone() - # TODO: handle a special version of CLIP with hooked_patches - k1 = () + new_clip = CLIPWithHooks(clip) + k1 = new_clip.add_hooked_patches(lora_hook=lora_hook, patches=loaded, strength_patch=strength_clip) else: k1 = () new_clip = None diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 0917a40..be7b9fc 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -6,7 +6,7 @@ from .nodes_gen1 import (AnimateDiffLoaderGen1, LegacyAnimateDiffLoaderWithConte from .nodes_gen2 import (UseEvolvedSamplingNode, ApplyAnimateDiffModelNode, ApplyAnimateDiffModelBasicNode, ApplyAnimateLCMI2VModel, ADKeyframeNode, LoadAnimateDiffModelNode, LoadAnimateLCMI2VModelNode, LoadAnimateDiffAndInjectI2VNode, UpscaleAndVaeEncode) from .nodes_multival import MultivalDynamicNode, MultivalScaledMaskNode -from .nodes_conditioning import MaskableLoraLoaderModelOnly, AttachLoraHook, CombineLoraHooks +from .nodes_conditioning import MaskableLoraLoader, MaskableLoraLoaderModelOnly, SetModelLoraHook, SetClipLoraHook, CombineLoraHooks from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGKeyframeNode) from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode) @@ -18,7 +18,7 @@ from .nodes_ad_settings import (AnimateDiffSettingsNode, ManualAdjustPENode, Swe from .nodes_extras import AnimateDiffUnload, EmptyLatentImageLarge, CheckpointLoaderSimpleWithNoiseSelect from .nodes_deprecated import (AnimateDiffLoader_Deprecated, AnimateDiffLoaderAdvanced_Deprecated, AnimateDiffCombine_Deprecated, AnimateDiffModelSettings, AnimateDiffModelSettingsSimple, AnimateDiffModelSettingsAdvanced, AnimateDiffModelSettingsAdvancedAttnStrengths) -from .nodes_lora import AnimateDiffLoraLoader, MaskedLoraLoader +from .nodes_lora import AnimateDiffLoraLoader from .logger import logger @@ -50,9 +50,11 @@ NODE_CLASS_MAPPINGS = { "ADE_IterationOptsDefault": IterationOptionsNode, "ADE_IterationOptsFreeInit": FreeInitOptionsNode, # Conditioning + "ADE_RegisterLoraHook": MaskableLoraLoader, "ADE_RegisterLoraHookModelOnly": MaskableLoraLoaderModelOnly, "ADE_CombineLoraHooks": CombineLoraHooks, - "ADE_AttachLoraHookToConditioning": AttachLoraHook, + "ADE_AttachLoraHookToConditioning": SetModelLoraHook, + "ADE_AttachLoraHookToCLIP": SetClipLoraHook, # Noise Layer Nodes "ADE_NoiseLayerAdd": NoiseLayerAddNode, "ADE_NoiseLayerAddWeighted": NoiseLayerAddWeightedNode, @@ -97,8 +99,6 @@ NODE_CLASS_MAPPINGS = { "ADE_LoadAnimateLCMI2VModel": LoadAnimateLCMI2VModelNode, "ADE_UpscaleAndVAEEncode": UpscaleAndVaeEncode, "ADE_InjectI2VIntoAnimateDiffModel": LoadAnimateDiffAndInjectI2VNode, - # MaskedLoraLoader - #"ADE_MaskedLoadLora": MaskedLoraLoader, # Deprecated Nodes "AnimateDiffLoaderV1": AnimateDiffLoader_Deprecated, "ADE_AnimateDiffLoaderV1Advanced": AnimateDiffLoaderAdvanced_Deprecated, @@ -127,9 +127,11 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_IterationOptsDefault": "Default Iteration Options πŸŽ­πŸ…πŸ…“", "ADE_IterationOptsFreeInit": "FreeInit Iteration Options πŸŽ­πŸ…πŸ…“", # Conditioning + "ADE_RegisterLoraHook": "Register LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_RegisterLoraHookModelOnly": "Register LoRA Hook (Model Only) πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooks": "Combine LoRA Hooks πŸŽ­πŸ…πŸ…“", - "ADE_AttachLoraHookToConditioning": "Attach LoRA Hook to Conditioning πŸŽ­πŸ…πŸ…“", + "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", + "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", @@ -174,8 +176,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_LoadAnimateLCMI2VModel": "Load AnimateLCM-I2V Model πŸŽ­πŸ…πŸ…“β‘‘", "ADE_UpscaleAndVAEEncode": "Scale Ref Image and VAE Encode πŸŽ­πŸ…πŸ…“β‘‘", "ADE_InjectI2VIntoAnimateDiffModel": "πŸ§ͺInject I2V into AnimateDiff Model πŸŽ­πŸ…πŸ…“β‘‘", - # MaskedLoraLoader - #"ADE_MaskedLoadLora": "Load LoRA (Masked) πŸŽ­πŸ…πŸ…“", # Deprecated Nodes "AnimateDiffLoaderV1": "🚫AnimateDiff Loader [DEPRECATED] πŸŽ­πŸ…πŸ…“", "ADE_AnimateDiffLoaderV1Advanced": "🚫AnimateDiff Loader (Advanced) [DEPRECATED] πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index bee0c28..32fb82a 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -7,7 +7,7 @@ from comfy.sd import CLIP import comfy.utils from .utils_motion import LoraHook, LoraHookGroup -from .model_injection import ModelPatcherAndInjector, load_hooked_lora_for_models +from .model_injection import ModelPatcherAndInjector, CLIPWithHooks, load_hooked_lora_for_models # based on ComfyUI's nodes.py LoraLoader class MaskableLoraLoader: @@ -77,7 +77,7 @@ class MaskableLoraLoaderModelOnly(MaskableLoraLoader): return (model_lora, lora_hook) -class AttachLoraHook: +class SetModelLoraHook: @classmethod def INPUT_TYPES(s): return { @@ -98,6 +98,27 @@ class AttachLoraHook: n[1]["lora_hook"] = lora_hook c.append(n) return (c, ) + + +class SetClipLoraHook: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "clip": ("CLIP",), + "lora_hook": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("CLIP",) + RETURN_NAMES = ("hook_CLIP",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "apply_lora_hook" + + def apply_lora_hook(self, clip: CLIP, lora_hook: LoraHookGroup): + new_clip = CLIPWithHooks(clip) + new_clip.set_desired_hooks(lora_hooks=lora_hook) + return (new_clip, ) class CombineLoraHooks: diff --git a/animatediff/nodes_lora.py b/animatediff/nodes_lora.py index a3db3ba..3286962 100644 --- a/animatediff/nodes_lora.py +++ b/animatediff/nodes_lora.py @@ -40,51 +40,3 @@ class AnimateDiffLoraLoader: prev_motion_lora.add_lora(lora_info) return (prev_motion_lora,) - - -class MaskedLoraLoader: - def __init__(self): - self.loaded_lora = None - - @classmethod - def INPUT_TYPES(s): - return {"required": { "model": ("MODEL",), - "clip": ("CLIP", ), - "lora_name": (folder_paths.get_filename_list("loras"), ), - "strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), - "strength_clip": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), - }} - #RETURN_TYPES = () - RETURN_TYPES = ("MODEL", "CLIP") - FUNCTION = "load_lora" - - CATEGORY = "loaders" - - def load_lora(self, model, clip, lora_name, strength_model, strength_clip): - if strength_model == 0 and strength_clip == 0: - return (model, clip) - - lora_path = folder_paths.get_full_path("loras", lora_name) - lora = None - if self.loaded_lora is not None: - if self.loaded_lora[0] == lora_path: - lora = self.loaded_lora[1] - else: - temp = self.loaded_lora - self.loaded_lora = None - del temp - - if lora is None: - lora = comfy.utils.load_torch_file(lora_path, safe_load=True) - self.loaded_lora = (lora_path, lora) - - from pathlib import Path - with open(Path(__file__).parent.parent.parent / "sd_lora_keys.txt", "w") as lfile: - for key in lora: - lfile.write(f"{key}:\t{lora[key].size()}\n") - - model_lora, clip_lora = comfy.sd.load_lora_for_models(model, clip, lora, strength_model, strength_clip) - #return (model_lora, clip_lora) - return (model, clip) - - \ No newline at end of file