From 6961b592416cbc493ddba7461af4b63a5449579d Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 19 Mar 2024 12:51:04 -0500 Subject: [PATCH 01/30] Progress on maskable loras --- animatediff/conditioning.py | 67 +++++++++++++++++++++++++++++++++++++ animatediff/nodes_lora.py | 2 +- 2 files changed, 68 insertions(+), 1 deletion(-) create mode 100644 animatediff/conditioning.py diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py new file mode 100644 index 0000000..b9da961 --- /dev/null +++ b/animatediff/conditioning.py @@ -0,0 +1,67 @@ +import folder_paths + +from comfy.model_patcher import ModelPatcher +import comfy.utils + +from .model_injection import ModelPatcherAndInjector + +# based on ComfyUI's nodes.py LoraLoader +class MaskableLoraLoader: + 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 = ("MODEL", "CLIP", "LORA_IDS") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "load_lora" + + def load_lora(self, model: ModelPatcher, clip: ModelPatcher, lora_name: str, strength_model: float, strength_clip: float): + 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) + + model_lora, clip_lora = model.clone(), clip.clone() + return (model_lora, clip_lora) + + +class MaskableLoraLoaderModelOnly(MaskableLoraLoader): + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "lora_name": (folder_paths.get_filename_list("loras"), ), + "strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), + } + } + + RETURN_TYPES = ("MODEL", "LORA_IDS") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "load_lora_model_only" + + def load_lora_model_only(self, model: ModelPatcher, lora_name: str, strength_model: float): + return (self.load_lora(model, None, lora_name, strength_model, 0)[0],) diff --git a/animatediff/nodes_lora.py b/animatediff/nodes_lora.py index 5cc3ba3..a3db3ba 100644 --- a/animatediff/nodes_lora.py +++ b/animatediff/nodes_lora.py @@ -83,7 +83,7 @@ class MaskedLoraLoader: 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) + 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) From 305e562f74f58c01b120ecb8f40f0b63f2b0944b Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 20 Mar 2024 14:21:34 -0500 Subject: [PATCH 02/30] Progress on lora masking --- animatediff/conditioning.py | 82 ++++++++----------------------- animatediff/model_injection.py | 7 +++ animatediff/nodes_conditioning.py | 67 +++++++++++++++++++++++++ 3 files changed, 95 insertions(+), 61 deletions(-) create mode 100644 animatediff/nodes_conditioning.py diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py index b9da961..6c657ca 100644 --- a/animatediff/conditioning.py +++ b/animatediff/conditioning.py @@ -1,67 +1,27 @@ -import folder_paths -from comfy.model_patcher import ModelPatcher -import comfy.utils -from .model_injection import ModelPatcherAndInjector - -# based on ComfyUI's nodes.py LoraLoader -class MaskableLoraLoader: +class LoraHookGroup: + ''' + Stores LoRA hooks to apply for conditioning + ''' 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}), - } - } + self.hooks = [] - RETURN_TYPES = ("MODEL", "CLIP", "LORA_IDS") - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" - FUNCTION = "load_lora" + def add(self, hook: str): + if hook not in self.hooks: + self.hooks.append(hook) + + def is_empty(self): + return len(self.hooks) == 0 - def load_lora(self, model: ModelPatcher, clip: ModelPatcher, lora_name: str, strength_model: float, strength_clip: float): - 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) + def clone(self): + cloned = LoraHookGroup() + for hook in self.hooks: + cloned.add(hook) + return cloned - model_lora, clip_lora = model.clone(), clip.clone() - return (model_lora, clip_lora) - - -class MaskableLoraLoaderModelOnly(MaskableLoraLoader): - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "model": ("MODEL",), - "lora_name": (folder_paths.get_filename_list("loras"), ), - "strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), - } - } - - RETURN_TYPES = ("MODEL", "LORA_IDS") - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" - FUNCTION = "load_lora_model_only" - - def load_lora_model_only(self, model: ModelPatcher, lora_name: str, strength_model: float): - return (self.load_lora(model, None, lora_name, strength_model, 0)[0],) + def clone_and_combine(self, other: 'LoraHookGroup'): + cloned = self.clone() + for hook in other.hooks: + cloned.add(hook) + return cloned diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 28e0643..0050c96 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -89,6 +89,13 @@ class ModelPatcherAndInjector(ModelPatcher): cloned.sample_settings = self.sample_settings cloned.motion_injection_params = self.motion_injection_params.clone() if self.motion_injection_params else self.motion_injection_params return cloned + + @classmethod + def create_from(cls, model: Union[ModelPatcher, 'ModelPatcherAndInjector']) -> 'ModelPatcherAndInjector': + if isinstance(model, ModelPatcherAndInjector): + return model.clone() + else: + return ModelPatcherAndInjector(model) class MotionModelPatcher(ModelPatcher): diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py new file mode 100644 index 0000000..b9da961 --- /dev/null +++ b/animatediff/nodes_conditioning.py @@ -0,0 +1,67 @@ +import folder_paths + +from comfy.model_patcher import ModelPatcher +import comfy.utils + +from .model_injection import ModelPatcherAndInjector + +# based on ComfyUI's nodes.py LoraLoader +class MaskableLoraLoader: + 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 = ("MODEL", "CLIP", "LORA_IDS") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "load_lora" + + def load_lora(self, model: ModelPatcher, clip: ModelPatcher, lora_name: str, strength_model: float, strength_clip: float): + 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) + + model_lora, clip_lora = model.clone(), clip.clone() + return (model_lora, clip_lora) + + +class MaskableLoraLoaderModelOnly(MaskableLoraLoader): + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "lora_name": (folder_paths.get_filename_list("loras"), ), + "strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), + } + } + + RETURN_TYPES = ("MODEL", "LORA_IDS") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "load_lora_model_only" + + def load_lora_model_only(self, model: ModelPatcher, lora_name: str, strength_model: float): + return (self.load_lora(model, None, lora_name, strength_model, 0)[0],) From de9be7b376444dbe2e661936535ea8218080997d Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 21 Mar 2024 19:27:11 -0500 Subject: [PATCH 03/30] Working proof-of-concept for LoRA masking --- animatediff/conditioning.py | 27 ----- animatediff/model_injection.py | 180 ++++++++++++++++++++++++++++-- animatediff/nodes.py | 9 ++ animatediff/nodes_conditioning.py | 65 +++++++++-- animatediff/nodes_deprecated.py | 4 +- animatediff/nodes_gen1.py | 4 +- animatediff/nodes_gen2.py | 2 +- animatediff/sampling.py | 145 +++++++++++++++++++++++- animatediff/utils_motion.py | 37 ++++++ 9 files changed, 422 insertions(+), 51 deletions(-) diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py index 6c657ca..e69de29 100644 --- a/animatediff/conditioning.py +++ b/animatediff/conditioning.py @@ -1,27 +0,0 @@ - - -class LoraHookGroup: - ''' - Stores LoRA hooks to apply for conditioning - ''' - def __init__(self): - self.hooks = [] - - def add(self, hook: str): - if hook not in self.hooks: - self.hooks.append(hook) - - def is_empty(self): - return len(self.hooks) == 0 - - def clone(self): - cloned = LoraHookGroup() - for hook in self.hooks: - cloned.add(hook) - return cloned - - def clone_and_combine(self, other: 'LoraHookGroup'): - cloned = self.clone() - for hook in other.hooks: - cloned.add(hook) - return cloned diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 4a67733..defdb3c 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -5,17 +5,20 @@ from einops import rearrange from torch import Tensor import torch.nn.functional as F import torch +import uuid import comfy.model_management import comfy.utils from comfy.model_patcher import ModelPatcher from comfy.model_base import BaseModel +from comfy.sd import CLIP from .ad_settings import AnimateDiffSettings, AdjustPE, AdjustWeight from .context import ContextOptions, ContextOptions, ContextOptionsGroup from .motion_module_ad import AnimateDiffModel, AnimateDiffFormat, EncoderOnlyAnimateDiffModel, has_mid_block, normalize_ad_state_dict from .logger import logger -from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max +from .utils_motion import (ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max, + LoraHook, LoraHookGroup) from .motion_lora import MotionLoraInfo, MotionLoraList from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type from .sample_settings import SampleSettings, SeedNoiseGeneration @@ -41,12 +44,48 @@ class ModelPatcherAndInjector(ModelPatcher): if hasattr(m, "object_patches_backup"): self.object_patches_backup = m.object_patches_backup - + # lora hook stuff + self.hooked_patches = {} + self.hooked_backup = {} + self.current_lora_hooks = None # injection stuff - self.motion_injection_params: InjectionParams = None + self.motion_injection_params: InjectionParams = InjectionParams() self.sample_settings: SampleSettings = SampleSettings() self.motion_models: MotionModelGroup = None + def add_hooked_patches(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): + ''' + Based on add_patches, but for hooked weights. + ''' + # TODO: make this work with timestep scheduling + current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {}) + p = set() + for key in patches: + if key in self.model_keys: + p.add(key) + current_patches: list[tuple] = current_hooked_patches.get(key, []) + current_patches.append((strength_patch, patches[key], strength_model)) + current_hooked_patches[key] = current_patches + self.hooked_patches[lora_hook] = current_hooked_patches + # since should care about these patches too to determine if same model, roll patches_uuid + self.patches_uuid = uuid.uuid4() + return list(p) + + def get_combined_hooked_patches(self, lora_hooks: LoraHookGroup): + ''' + Returns patches for selected lora_hooks. + ''' + # combined_patches will contain weights of all relevant lora_hooks, per key + combined_patches = {} + if lora_hooks is not None: + for hook in lora_hooks.hooks: + hook_patches: dict = self.hooked_patches.get(hook, {}) + for key in hook_patches.keys(): + current_patches: list[tuple] = combined_patches.get(key, []) + current_patches.extend(hook_patches[key]) + combined_patches[key] = current_patches + return combined_patches + def model_patches_to(self, device): super().model_patches_to(device) if self.motion_models is not None: @@ -71,6 +110,8 @@ class ModelPatcherAndInjector(ModelPatcher): self.eject_model(device_to=device_to) # finally, do normal model unpatching if unpatch_weights: # TODO: keep only 'else' portion when don't need to worry about past comfy versions + # handle hooked_patches first + self.unpatch_hooked(device_to=device_to) return super().unpatch_model(device_to) else: return super().unpatch_model(device_to, unpatch_weights) @@ -93,21 +134,141 @@ class ModelPatcherAndInjector(ModelPatcher): except Exception: pass - def clone(self): + def apply_lora_hooks(self, lora_hooks: LoraHookGroup, device_to=None): + # first, determine if need to reapply patches + if self.current_lora_hooks == lora_hooks: + return + # unpatch hooks, if needed + self.unpatch_hooked(device_to=device_to) + # finally, patch hooks + # TODO: handle lowvram + self.patch_hooked(lora_hooks=lora_hooks, device_to=device_to) + + def patch_hooked(self, lora_hooks: LoraHookGroup, device_to=None, patch_weights=True) -> None: + if not patch_weights: + return + # use current device + if not device_to: + device_to=self.current_device + # first, unpatch any previous patches + self.unpatch_hooked() + # then, handle weights + if patch_weights: + model_sd = self.model_state_dict() + # get combined patches of relevant lora_hooks + relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) + for key in relevant_patches: + if key not in model_sd: + logger.warning(f"LoraHook hook could not patch. key doesn't exist in model: {key}") + continue + self.patch_hooked_weight_to_device(combined_patches=relevant_patches, key=key, device_to=device_to) + self.current_lora_hooks = lora_hooks + + def patch_hooked_lowvram(self, lora_hooks: LoraHookGroup, device_to=None, lowvram_model_memory=0): + # TODO: handle lowvram situation + pass + + def patch_hooked_weight_to_device(self, combined_patches: dict, key: str, device_to=None): + if key not in combined_patches: + return + + weight: Tensor = comfy.utils.get_attr(self.model, key) + + inplace_update = self.weight_inplace_update + + if key not in self.hooked_backup: + self.hooked_backup[key] = weight.to(device=self.offload_device, copy=inplace_update) + + if device_to is not None: + temp_weight = comfy.model_management.cast_to_device(weight, device_to, torch.float32, copy=True) + else: + temp_weight = weight.to(torch.float32, copy=True) + out_weight = self.calculate_weight(combined_patches[key], temp_weight, key).to(weight.dtype) + if inplace_update: + comfy.utils.copy_to_param(self.model, key, out_weight) + else: + comfy.utils.set_attr_param(self.model, key, out_weight) + + def unpatch_hooked(self, device_to=None, unpatch_weights=True) -> None: + if not unpatch_weights: + return + # if no backups from before hook, then nothing to unpatch + if len(self.hooked_backup) == 0: + return + # TODO: handle lowvram, assuming there is something that needs to be done + if self.model_lowvram: + pass + keys = list(self.hooked_backup.keys()) + + if self.weight_inplace_update: + for k in keys: + if device_to is None: + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) + else: + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k]) + else: + for k in keys: + if device_to is None: + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) + else: + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k]) + # clear hooked_backup + self.hooked_backup.clear() + self.current_lora_hooks = None + + def clone(self, hooks_only=False): cloned = ModelPatcherAndInjector(self) - cloned.motion_models = self.motion_models.clone() if self.motion_models else self.motion_models - cloned.sample_settings = self.sample_settings - cloned.motion_injection_params = self.motion_injection_params.clone() if self.motion_injection_params else self.motion_injection_params + for hook in self.hooked_patches: + cloned.hooked_patches[hook] = {} + for k in self.hooked_patches[hook]: + cloned.hooked_patches[hook][k] = self.hooked_patches[hook][k][:] + cloned.hooked_backup = self.hooked_backup + cloned.current_lora_hooks = self.current_lora_hooks + if not hooks_only: + cloned.motion_models = self.motion_models.clone() if self.motion_models else self.motion_models + cloned.sample_settings = self.sample_settings + cloned.motion_injection_params = self.motion_injection_params.clone() if self.motion_injection_params else self.motion_injection_params return cloned @classmethod - def create_from(cls, model: Union[ModelPatcher, 'ModelPatcherAndInjector']) -> 'ModelPatcherAndInjector': + def create_from(cls, model: Union[ModelPatcher, 'ModelPatcherAndInjector'], hooks_only=False) -> 'ModelPatcherAndInjector': if isinstance(model, ModelPatcherAndInjector): - return model.clone() + return model.clone(hooks_only=hooks_only) else: return ModelPatcherAndInjector(model) +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: + key_map = comfy.lora.model_lora_keys_unet(model.model, key_map) + if clip is not None: + key_map = comfy.lora.model_lora_keys_clip(clip.cond_stage_model, key_map) + + loaded = comfy.lora.load_lora(lora, key_map) + if model is not None: + new_modelpatcher = ModelPatcherAndInjector.create_from(model) + k = new_modelpatcher.add_hooked_patches(lora_hook=lora_hook, patches=loaded, strength_patch=strength_model) + else: + k = () + new_modelpatcher = None + + if clip is not None: + new_clip = clip.clone() + # TODO: handle a special version of CLIP with hooked_patches + k1 = () + else: + k1 = () + new_clip = None + k = set(k) + k1 = set(k1) + for x in loaded: + if (x not in k) and (x not in k1): + logger.warning(f"NOT LOADED {x}") + + return (new_modelpatcher, new_clip) + + class MotionModelPatcher(ModelPatcher): # Mostly here so that type hints work in IDEs def __init__(self, *args, **kwargs): @@ -359,6 +520,7 @@ def get_vanilla_model_patcher(m: ModelPatcher) -> ModelPatcher: model.model_keys = m.model_keys return model + # 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 diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 09535f5..0917a40 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -6,6 +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_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGKeyframeNode) from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode) @@ -48,6 +49,10 @@ NODE_CLASS_MAPPINGS = { # Iteration Opts "ADE_IterationOptsDefault": IterationOptionsNode, "ADE_IterationOptsFreeInit": FreeInitOptionsNode, + # Conditioning + "ADE_RegisterLoraHookModelOnly": MaskableLoraLoaderModelOnly, + "ADE_CombineLoraHooks": CombineLoraHooks, + "ADE_AttachLoraHookToConditioning": AttachLoraHook, # Noise Layer Nodes "ADE_NoiseLayerAdd": NoiseLayerAddNode, "ADE_NoiseLayerAddWeighted": NoiseLayerAddWeightedNode, @@ -121,6 +126,10 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Iteration Opts "ADE_IterationOptsDefault": "Default Iteration Options πŸŽ­πŸ…πŸ…“", "ADE_IterationOptsFreeInit": "FreeInit Iteration Options πŸŽ­πŸ…πŸ…“", + # Conditioning + "ADE_RegisterLoraHookModelOnly": "Register LoRA Hook (Model Only) πŸŽ­πŸ…πŸ…“", + "ADE_CombineLoraHooks": "Combine LoRA Hooks πŸŽ­πŸ…πŸ…“", + "ADE_AttachLoraHookToConditioning": "Attach LoRA Hook to Conditioning πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index b9da961..bee0c28 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -1,9 +1,13 @@ +import uuid import folder_paths +from typing import Union from comfy.model_patcher import ModelPatcher +from comfy.sd import CLIP import comfy.utils -from .model_injection import ModelPatcherAndInjector +from .utils_motion import LoraHook, LoraHookGroup +from .model_injection import ModelPatcherAndInjector, load_hooked_lora_for_models # based on ComfyUI's nodes.py LoraLoader class MaskableLoraLoader: @@ -22,11 +26,11 @@ class MaskableLoraLoader: } } - RETURN_TYPES = ("MODEL", "CLIP", "LORA_IDS") + RETURN_TYPES = ("MODEL", "CLIP", "LORA_HOOK") CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" FUNCTION = "load_lora" - def load_lora(self, model: ModelPatcher, clip: ModelPatcher, lora_name: str, strength_model: float, strength_clip: float): + def load_lora(self, model: Union[ModelPatcher, ModelPatcherAndInjector], clip: CLIP, lora_name: str, strength_model: float, strength_clip: float): if strength_model == 0 and strength_clip == 0: return (model, clip) @@ -44,8 +48,12 @@ class MaskableLoraLoader: lora = comfy.utils.load_torch_file(lora_path, safe_load=True) self.loaded_lora = (lora_path, lora) - model_lora, clip_lora = model.clone(), clip.clone() - return (model_lora, clip_lora) + lora_hook = LoraHook(lora_name=lora_name) + lora_hook_group = LoraHookGroup() + lora_hook_group.add(lora_hook) + model_lora, clip_lora = load_hooked_lora_for_models(model=model, clip=clip, lora=lora, lora_hook=lora_hook, + strength_model=strength_model, strength_clip=strength_clip) + return (model_lora, clip_lora, lora_hook_group) class MaskableLoraLoaderModelOnly(MaskableLoraLoader): @@ -59,9 +67,52 @@ class MaskableLoraLoaderModelOnly(MaskableLoraLoader): } } - RETURN_TYPES = ("MODEL", "LORA_IDS") + RETURN_TYPES = ("MODEL", "LORA_HOOK") CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" FUNCTION = "load_lora_model_only" def load_lora_model_only(self, model: ModelPatcher, lora_name: str, strength_model: float): - return (self.load_lora(model, None, lora_name, strength_model, 0)[0],) + model_lora, clip_lora, lora_hook = self.load_lora(model=model, clip=None, lora_name=lora_name, + strength_model=strength_model, strength_clip=0) + return (model_lora, lora_hook) + + +class AttachLoraHook: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "conditioning": ("CONDITIONING",), + "lora_hook": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("CONDITIONING",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "attach_lora_hook" + + def attach_lora_hook(self, conditioning, lora_hook: LoraHookGroup): + c = [] + for t in conditioning: + n = [t[0], t[1].copy()] + n[1]["lora_hook"] = lora_hook + c.append(n) + return (c, ) + + +class CombineLoraHooks: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "lora_hook_A": ("LORA_HOOK",), + "lora_hook_B": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("LORA_HOOK",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "combine_lora_hooks" + + def combine_lora_hooks(self, lora_hook_A: LoraHookGroup, lora_hook_B: LoraHookGroup): + return (lora_hook_A.clone_and_combine(lora_hook_B),) diff --git a/animatediff/nodes_deprecated.py b/animatediff/nodes_deprecated.py index c7349bc..3e8feb1 100644 --- a/animatediff/nodes_deprecated.py +++ b/animatediff/nodes_deprecated.py @@ -54,7 +54,7 @@ class AnimateDiffLoader_Deprecated: apply_v2_properly=False, ) # inject for use in sampling code - model = ModelPatcherAndInjector(model) + model = ModelPatcherAndInjector.create_from(model, hooks_only=True) model.motion_models = MotionModelGroup(motion_model) model.motion_injection_params = params @@ -123,7 +123,7 @@ class AnimateDiffLoaderAdvanced_Deprecated: # set context settings params.set_context(context_options=context_group) # inject for use in sampling code - model = ModelPatcherAndInjector(model) + model = ModelPatcherAndInjector.create_from(model, hooks_only=True) model.motion_models = MotionModelGroup(motion_model) model.motion_injection_params = params diff --git a/animatediff/nodes_gen1.py b/animatediff/nodes_gen1.py index 6204514..ad93440 100644 --- a/animatediff/nodes_gen1.py +++ b/animatediff/nodes_gen1.py @@ -76,7 +76,7 @@ class AnimateDiffLoaderGen1: # need to use a ModelPatcher that supports injection of motion modules into unet # need to use a ModelPatcher that supports injection of motion modules into unet - model = ModelPatcherAndInjector(model) + model = ModelPatcherAndInjector.create_from(model, hooks_only=True) model.motion_models = MotionModelGroup(motion_model) model.sample_settings = sample_settings if sample_settings is not None else SampleSettings() model.motion_injection_params = params @@ -157,7 +157,7 @@ class LegacyAnimateDiffLoaderWithContext: motion_model.keyframes = ad_keyframes.clone() if ad_keyframes else ADKeyframeGroup() - model = ModelPatcherAndInjector(model) + model = ModelPatcherAndInjector.create_from(model, hooks_only=True) model.motion_models = MotionModelGroup(motion_model) model.sample_settings = sample_settings if sample_settings is not None else SampleSettings() model.motion_injection_params = params diff --git a/animatediff/nodes_gen2.py b/animatediff/nodes_gen2.py index c7829a5..9bac59e 100644 --- a/animatediff/nodes_gen2.py +++ b/animatediff/nodes_gen2.py @@ -57,7 +57,7 @@ class UseEvolvedSamplingNode: if context_options: params.set_context(context_options) # need to use a ModelPatcher that supports injection of motion modules into unet - model = ModelPatcherAndInjector(model) + model = ModelPatcherAndInjector.create_from(model, hooks_only=True) model.motion_models = m_models model.sample_settings = sample_settings if sample_settings is not None else SampleSettings() model.motion_injection_params = params diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 25a2e1e..76ceae5 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -18,7 +18,7 @@ import comfy.ops from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows from .sample_settings import IterationOptions, SampleSettings, SeedNoiseGeneration, prepare_mask_ad from .utils_model import ModelTypeSD, wrap_function_to_inject_xformers_bug_info -from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher +from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, LoraHookGroup from .motion_module_ad import AnimateDiffFormat, AnimateDiffInfo, AnimateDiffVersion, VanillaTemporalModule from .logger import logger @@ -28,6 +28,7 @@ from .logger import logger # Global variable to use to more conveniently hack variable access into samplers class AnimateDiffHelper_GlobalState: def __init__(self): + self.model_patcher: ModelPatcherAndInjector = None self.motion_models: MotionModelGroup = None self.params: InjectionParams = None self.sample_settings: SampleSettings = None @@ -50,6 +51,9 @@ class AnimateDiffHelper_GlobalState: self.last_step: int = 0 self.current_step: int = 0 self.total_steps: int = 0 + if self.model_patcher is not None: + del self.model_patcher + self.model_patcher = None if self.motion_models is not None: del self.motion_models self.motion_models = None @@ -306,6 +310,7 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) -> # update GLOBALSTATE for next iteration ADGS.current_step = ADGS.start_step + step + 1 kwargs["callback"] = ad_callback + ADGS.model_patcher = model ADGS.motion_models = model.motion_models ADGS.sample_settings = model.sample_settings @@ -402,7 +407,7 @@ def evolved_sampling_function(model, x, timestep, uncond, cond, cond_scale, mode model_options["transformer_options"]["ad_params"] = ADGS.create_exposed_params() if not ADGS.is_using_sliding_context(): - cond_pred, uncond_pred = comfy.samplers.calc_cond_uncond_batch(model, cond, uncond_, x, timestep, model_options) + cond_pred, uncond_pred = calc_cond_uncond_batch_wrapper(model, cond, uncond_, x, timestep, model_options) else: cond_pred, uncond_pred = sliding_calc_cond_uncond_batch(model, cond, uncond_, x, timestep, model_options) @@ -516,7 +521,7 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, sub_cond = get_resized_cond(cond, full_idxs, len(ctx_idxs)) if cond is not None else None sub_uncond = get_resized_cond(uncond, full_idxs, len(ctx_idxs)) if uncond is not None else None - sub_cond_out, sub_uncond_out = comfy.samplers.calc_cond_uncond_batch(model, sub_cond, sub_uncond, sub_x, sub_timestep, model_options) + sub_cond_out, sub_uncond_out = calc_cond_uncond_batch_wrapper(model, sub_cond, sub_uncond, sub_x, sub_timestep, model_options) if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE: full_length = ADGS.params.full_length @@ -551,3 +556,137 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, uncond_final /= out_count_final del out_count_final return cond_final, uncond_final + + +def calc_cond_uncond_batch_wrapper(model, cond, uncond, x_in, timestep, model_options): + # check if conds or unconds contain lora_hook + contains_lora_hooks = False + for cond_uncond in [cond, uncond]: + for t in cond_uncond: + if "lora_hook" in t: + contains_lora_hooks = True + break + if contains_lora_hooks: + break + if contains_lora_hooks: + return calc_cond_uncond_batch_lora_hook(model, cond, uncond, x_in, timestep, model_options) + return comfy.samplers.calc_cond_uncond_batch(model, cond, uncond, x_in, timestep, model_options) + + +def calc_cond_uncond_batch_lora_hook(model, cond, uncond, x_in, timestep, model_options): + out_cond = torch.zeros_like(x_in) + out_count = torch.ones_like(x_in) * 1e-37 + + out_uncond = torch.zeros_like(x_in) + out_uncond_count = torch.ones_like(x_in) * 1e-37 + + COND = 0 + UNCOND = 1 + + # separate conds and unconds by matching lora_hooks + hooked_to_run = {} + for x in cond: + p = comfy.samplers.get_area_and_mult(x, x_in, timestep) + if p is None: + continue + hook: LoraHookGroup = x.get("lora_hook", None) + hooked_to_run.setdefault(hook, list()) + hooked_to_run[hook] += [(p, COND)] + if uncond is not None: + for x in uncond: + p = comfy.samplers.get_area_and_mult(x, x_in, timestep) + if p is None: + continue + hook: LoraHookGroup = x.get("lora_hook", None) + hooked_to_run.setdefault(hook, list()) + hooked_to_run[hook] += [(p, UNCOND)] + + # run every hooked_to_run separately + for lora_hooks, to_run in hooked_to_run.items(): + while len(to_run) > 0: + first = to_run[0] + first_shape = first[0][0].shape + to_batch_temp = [] + for x in range(len(to_run)): + if comfy.samplers.can_concat_cond(to_run[x][0], first[0]): + to_batch_temp += [x] + + to_batch_temp.reverse() + to_batch = to_batch_temp[:1] + + free_memory = model_management.get_free_memory(x_in.device) + for i in range(1, len(to_batch_temp) + 1): + batch_amount = to_batch_temp[:len(to_batch_temp)//i] + input_shape = [len(batch_amount) * first_shape[0]] + list(first_shape)[1:] + if model.memory_required(input_shape) < free_memory: + to_batch = batch_amount + break + ADGS.model_patcher.apply_lora_hooks(lora_hooks=lora_hooks) + + input_x = [] + mult = [] + c = [] + cond_or_uncond = [] + area = [] + control = None + patches = None + for x in to_batch: + o = to_run.pop(x) + p = o[0] + input_x.append(p.input_x) + mult.append(p.mult) + c.append(p.conditioning) + area.append(p.area) + cond_or_uncond.append(o[1]) + control = p.control + patches = p.patches + + batch_chunks = len(cond_or_uncond) + input_x = torch.cat(input_x) + c = comfy.samplers.cond_cat(c) + timestep_ = torch.cat([timestep] * batch_chunks) + + if control is not None: + c['control'] = control.get_control(input_x, timestep_, c, len(cond_or_uncond)) + + transformer_options = {} + if 'transformer_options' in model_options: + transformer_options = model_options['transformer_options'].copy() + + if patches is not None: + if "patches" in transformer_options: + cur_patches = transformer_options["patches"].copy() + for p in patches: + if p in cur_patches: + cur_patches[p] = cur_patches[p] + patches[p] + else: + cur_patches[p] = patches[p] + transformer_options["patches"] = cur_patches + else: + transformer_options["patches"] = patches + + transformer_options["cond_or_uncond"] = cond_or_uncond[:] + transformer_options["sigmas"] = timestep + + c['transformer_options'] = transformer_options + + if 'model_function_wrapper' in model_options: + output = model_options['model_function_wrapper'](model.apply_model, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond}).chunk(batch_chunks) + else: + output = model.apply_model(input_x, timestep_, **c).chunk(batch_chunks) + del input_x + + for o in range(batch_chunks): + if cond_or_uncond[o] == COND: + out_cond[:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += output[o] * mult[o] + out_count[:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += mult[o] + else: + out_uncond[:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += output[o] * mult[o] + out_uncond_count[:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += mult[o] + del mult + + out_cond /= out_count + del out_count + out_uncond /= out_uncond_count + del out_uncond_count + return out_cond, out_uncond diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index 0aaac63..11ca50a 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -2,6 +2,7 @@ from typing import Union import torch import torch.nn.functional as F from torch import Tensor, nn +import uuid import comfy.model_management as model_management import comfy.ops @@ -250,6 +251,42 @@ class ADKeyframeGroup: return cloned +class LoraHook: + def __init__(self, lora_name: str): + self.lora_name = lora_name + self.id = f"{lora_name}|{uuid.uuid4()}" + + # def __eq__(self, other: 'LoraHook'): + # return self.id == other.id + + +class LoraHookGroup: + ''' + Stores LoRA hooks to apply for conditioning + ''' + def __init__(self): + self.hooks = [] + + def add(self, hook: str): + if hook not in self.hooks: + self.hooks.append(hook) + + def is_empty(self): + return len(self.hooks) == 0 + + def clone(self): + cloned = LoraHookGroup() + for hook in self.hooks: + cloned.add(hook) + return cloned + + def clone_and_combine(self, other: 'LoraHookGroup'): + cloned = self.clone() + for hook in other.hooks: + cloned.add(hook) + return cloned + + class DummyNNModule(nn.Module): class DoNothingWhenCalled: def __call__(self, *args, **kwargs): From 448d7463341365e7a9e7e0aa86a998a1e0ec61e0 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 21 Mar 2024 23:41:33 -0500 Subject: [PATCH 04/30] Made sure ejection/injection happens between hooked_patches patching/unpatching (might be unnecessary, we'll see) --- animatediff/model_injection.py | 21 +++++++++++++++++++-- 1 file changed, 19 insertions(+), 2 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index defdb3c..371eaf5 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -49,6 +49,7 @@ class ModelPatcherAndInjector(ModelPatcher): self.hooked_backup = {} self.current_lora_hooks = None # injection stuff + self.currently_injected = False self.motion_injection_params: InjectionParams = InjectionParams() self.sample_settings: SampleSettings = SampleSettings() self.motion_models: MotionModelGroup = None @@ -119,6 +120,7 @@ class ModelPatcherAndInjector(ModelPatcher): def inject_model(self, device_to=None): if self.motion_models is not None: for motion_model in self.motion_models.models: + self.currently_injected = True motion_model.model.inject(self) try: motion_model.model.to(device_to) @@ -133,6 +135,7 @@ class ModelPatcherAndInjector(ModelPatcher): motion_model.model.to(device_to) except Exception: pass + self.currently_injected = False def apply_lora_hooks(self, lora_hooks: LoraHookGroup, device_to=None): # first, determine if need to reapply patches @@ -150,10 +153,13 @@ class ModelPatcherAndInjector(ModelPatcher): # use current device if not device_to: device_to=self.current_device - # first, unpatch any previous patches - self.unpatch_hooked() # then, handle weights if patch_weights: + # first, unpatch any previous patches + self.unpatch_hooked() + was_injected = self.currently_injected + if was_injected: + self.eject_model() model_sd = self.model_state_dict() # get combined patches of relevant lora_hooks relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) @@ -163,6 +169,10 @@ class ModelPatcherAndInjector(ModelPatcher): continue self.patch_hooked_weight_to_device(combined_patches=relevant_patches, key=key, device_to=device_to) self.current_lora_hooks = lora_hooks + # reinject model, if needed + if was_injected: + self.inject_model() + def patch_hooked_lowvram(self, lora_hooks: LoraHookGroup, device_to=None, lowvram_model_memory=0): # TODO: handle lowvram situation @@ -195,6 +205,9 @@ class ModelPatcherAndInjector(ModelPatcher): # if no backups from before hook, then nothing to unpatch if len(self.hooked_backup) == 0: return + was_injected = self.currently_injected + if was_injected: + self.eject_model() # TODO: handle lowvram, assuming there is something that needs to be done if self.model_lowvram: pass @@ -215,6 +228,9 @@ class ModelPatcherAndInjector(ModelPatcher): # clear hooked_backup self.hooked_backup.clear() self.current_lora_hooks = None + # reinject model, if necessary + if was_injected: + self.inject_model() def clone(self, hooks_only=False): cloned = ModelPatcherAndInjector(self) @@ -224,6 +240,7 @@ class ModelPatcherAndInjector(ModelPatcher): cloned.hooked_patches[hook][k] = self.hooked_patches[hook][k][:] cloned.hooked_backup = self.hooked_backup cloned.current_lora_hooks = self.current_lora_hooks + cloned.currently_injected = self.currently_injected if not hooks_only: cloned.motion_models = self.motion_models.clone() if self.motion_models else self.motion_models cloned.sample_settings = self.sample_settings From a9a0082d7cbcd1cbec74b79afac0462e9340c90c Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 22 Mar 2024 01:50:07 -0500 Subject: [PATCH 05/30] Added None check in calc_cond_uncond_batch_wrapper (uncond can be None) --- animatediff/sampling.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 76ceae5..bf577b2 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -562,6 +562,8 @@ def calc_cond_uncond_batch_wrapper(model, cond, uncond, x_in, timestep, model_op # check if conds or unconds contain lora_hook contains_lora_hooks = False for cond_uncond in [cond, uncond]: + if cond_uncond is None: + continue for t in cond_uncond: if "lora_hook" in t: contains_lora_hooks = True From 2ae255e97d29379518d293eca80d647609f3c14e Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 23 Mar 2024 12:12:01 -0500 Subject: [PATCH 06/30] 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 From 6f6b887b71ec5de49bba707757d1e8ae17488187 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 24 Mar 2024 23:17:47 -0500 Subject: [PATCH 07/30] Added MAX_SPEED lora masking mode, fixed masking for newest comfy updates --- animatediff/model_injection.py | 117 +++++++++++++++++++++++---------- animatediff/sample_settings.py | 2 +- animatediff/sampling.py | 4 ++ animatediff/utils_motion.py | 13 +++- 4 files changed, 99 insertions(+), 37 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index da0fe6d..4e07d76 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -18,7 +18,7 @@ from .context import ContextOptions, ContextOptions, ContextOptionsGroup from .motion_module_ad import AnimateDiffModel, AnimateDiffFormat, EncoderOnlyAnimateDiffModel, has_mid_block, normalize_ad_state_dict from .logger import logger from .utils_motion import (ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max, - LoraHook, LoraHookGroup) + LoraHook, LoraHookGroup, LoraHookMode) from .motion_lora import MotionLoraInfo, MotionLoraList from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type from .sample_settings import SampleSettings, SeedNoiseGeneration @@ -45,15 +45,49 @@ class ModelPatcherAndInjector(ModelPatcher): self.object_patches_backup = m.object_patches_backup # lora hook stuff - self.hooked_patches = {} + self.hooked_patches = {} # binds LoraHook to specific keys self.hooked_backup = {} + self.cached_hooked_patches = {} # binds LoraHookGroup to pre-calculated weights (speed optimization) self.current_lora_hooks = None + self.lora_hook_mode = LoraHookMode.MAX_SPEED # injection stuff self.currently_injected = False self.motion_injection_params: InjectionParams = InjectionParams() self.sample_settings: SampleSettings = SampleSettings() self.motion_models: MotionModelGroup = None + def clone(self, hooks_only=False): + cloned = ModelPatcherAndInjector(self) + # copy lora hooks + for hook in self.hooked_patches: + cloned.hooked_patches[hook] = {} + for k in self.hooked_patches[hook]: + cloned.hooked_patches[hook][k] = self.hooked_patches[hook][k][:] + # copy pre-calc weights bound to LoraHookGroups + for group in self.cached_hooked_patches: + cloned.cached_hooked_patches[group] = {} + for k in self.cached_hooked_patches[group]: + cloned.cached_hooked_patches[group][k] = self.cached_hooked_patches[group][k] + cloned.hooked_backup = self.hooked_backup + cloned.current_lora_hooks = self.current_lora_hooks + cloned.currently_injected = self.currently_injected + cloned.lora_hook_mode = self.lora_hook_mode + if not hooks_only: + cloned.motion_models = self.motion_models.clone() if self.motion_models else self.motion_models + cloned.sample_settings = self.sample_settings + cloned.motion_injection_params = self.motion_injection_params.clone() if self.motion_injection_params else self.motion_injection_params + return cloned + + @classmethod + def create_from(cls, model: Union[ModelPatcher, 'ModelPatcherAndInjector'], hooks_only=False) -> 'ModelPatcherAndInjector': + if isinstance(model, ModelPatcherAndInjector): + return model.clone(hooks_only=hooks_only) + else: + return ModelPatcherAndInjector(model) + + def set_lora_hook_mode(self, lora_hook_mode: str): + self.lora_hook_mode = lora_hook_mode + def add_hooked_patches(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): ''' Based on add_patches, but for hooked weights. @@ -113,6 +147,7 @@ class ModelPatcherAndInjector(ModelPatcher): if unpatch_weights: # TODO: keep only 'else' portion when don't need to worry about past comfy versions # handle hooked_patches first self.unpatch_hooked(device_to=device_to) + self.clear_cached_hooked_weights() return super().unpatch_model(device_to) else: return super().unpatch_model(device_to, unpatch_weights) @@ -157,43 +192,77 @@ class ModelPatcherAndInjector(ModelPatcher): if patch_weights: # first, unpatch any previous patches self.unpatch_hooked() + # eject model, if needed was_injected = self.currently_injected if was_injected: self.eject_model() + model_sd = self.model_state_dict() - # get combined patches of relevant lora_hooks - relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) - for key in relevant_patches: - if key not in model_sd: - logger.warning(f"LoraHook hook could not patch. key doesn't exist in model: {key}") - continue - self.patch_hooked_weight_to_device(combined_patches=relevant_patches, key=key, device_to=device_to) + # if have cached weights for lora_hooks, use it + cached_weights = self.cached_hooked_patches.get(lora_hooks, None) + if cached_weights is not None: + for key in cached_weights: + if key not in model_sd: + logger.warning(f"Cached LoraHook hook could not patch. key doesn't exist in model: {key}") + self.patch_cached_hooked_weight(cached_weights=cached_weights, key=key) + else: + # get combined patches of relevant lora_hooks + relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) + for key in relevant_patches: + if key not in model_sd: + logger.warning(f"LoraHook hook could not patch. key doesn't exist in model: {key}") + continue + self.patch_hooked_weight_to_device(lora_hooks=lora_hooks, combined_patches=relevant_patches, key=key, device_to=device_to) self.current_lora_hooks = lora_hooks # reinject model, if needed if was_injected: self.inject_model() - def patch_hooked_lowvram(self, lora_hooks: LoraHookGroup, device_to=None, lowvram_model_memory=0): # TODO: handle lowvram situation pass - def patch_hooked_weight_to_device(self, combined_patches: dict, key: str, device_to=None): + def patch_cached_hooked_weight(self, cached_weights: dict, key: str): + inplace_update = self.weight_inplace_update + target_device = self.offload_device + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + target_device = self.current_device + + weight: Tensor = comfy.utils.get_attr(self.model, key) + + if key not in self.hooked_backup: + self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) + if inplace_update: + comfy.utils.copy_to_param(self.model, key, cached_weights[key]) + else: + comfy.utils.set_attr_param(self.model, key, cached_weights[key]) + + def clear_cached_hooked_weights(self): + self.cached_hooked_patches.clear() + self.current_lora_hooks = None + + def patch_hooked_weight_to_device(self, lora_hooks: LoraHookGroup, combined_patches: dict, key: str, device_to=None): if key not in combined_patches: return weight: Tensor = comfy.utils.get_attr(self.model, key) inplace_update = self.weight_inplace_update + target_device = self.offload_device + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + target_device = self.current_device if key not in self.hooked_backup: - self.hooked_backup[key] = weight.to(device=self.offload_device, copy=inplace_update) + self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) if device_to is not None: temp_weight = comfy.model_management.cast_to_device(weight, device_to, torch.float32, copy=True) else: temp_weight = weight.to(torch.float32, copy=True) out_weight = self.calculate_weight(combined_patches[key], temp_weight, key).to(weight.dtype) + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + self.cached_hooked_patches.setdefault(lora_hooks, {}) + self.cached_hooked_patches[lora_hooks][key] = out_weight if inplace_update: comfy.utils.copy_to_param(self.model, key, out_weight) else: @@ -212,7 +281,6 @@ class ModelPatcherAndInjector(ModelPatcher): if self.model_lowvram: pass keys = list(self.hooked_backup.keys()) - if self.weight_inplace_update: for k in keys: if device_to is None: @@ -232,33 +300,12 @@ class ModelPatcherAndInjector(ModelPatcher): if was_injected: self.inject_model() - def clone(self, hooks_only=False): - cloned = ModelPatcherAndInjector(self) - for hook in self.hooked_patches: - cloned.hooked_patches[hook] = {} - for k in self.hooked_patches[hook]: - cloned.hooked_patches[hook][k] = self.hooked_patches[hook][k][:] - cloned.hooked_backup = self.hooked_backup - cloned.current_lora_hooks = self.current_lora_hooks - cloned.currently_injected = self.currently_injected - if not hooks_only: - cloned.motion_models = self.motion_models.clone() if self.motion_models else self.motion_models - cloned.sample_settings = self.sample_settings - cloned.motion_injection_params = self.motion_injection_params.clone() if self.motion_injection_params else self.motion_injection_params - return cloned - - @classmethod - def create_from(cls, model: Union[ModelPatcher, 'ModelPatcherAndInjector'], hooks_only=False) -> 'ModelPatcherAndInjector': - if isinstance(model, ModelPatcherAndInjector): - return model.clone(hooks_only=hooks_only) - else: - 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.patcher.set_lora_hook_mode(lora_hook_mode=LoraHookMode.MIN_VRAM) self.cond_stage_model = clip.cond_stage_model self.tokenizer = clip.tokenizer self.layer_idx = clip.layer_idx diff --git a/animatediff/sample_settings.py b/animatediff/sample_settings.py index 124b67b..c74c376 100644 --- a/animatediff/sample_settings.py +++ b/animatediff/sample_settings.py @@ -11,7 +11,7 @@ from comfy.model_base import BaseModel from . import freeinit from .context import ContextOptions, ContextOptionsGroup from .utils_model import SigmaSchedule -from .utils_motion import extend_to_batch_size, get_sorted_list_via_attr, prepare_mask_batch +from .utils_motion import extend_to_batch_size, get_sorted_list_via_attr, prepare_mask_batch, LoraHookMode from .logger import logger diff --git a/animatediff/sampling.py b/animatediff/sampling.py index bf577b2..e202179 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -52,6 +52,8 @@ class AnimateDiffHelper_GlobalState: self.current_step: int = 0 self.total_steps: int = 0 if self.model_patcher is not None: + self.model_patcher.unpatch_hooked() + self.model_patcher.clear_cached_hooked_weights() del self.model_patcher self.model_patcher = None if self.motion_models is not None: @@ -275,6 +277,8 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) -> cached_noise = None function_injections = FunctionInjectionHolder() try: + if len(model.hooked_patches) > 0: + model_management.cleanup_models() if model.sample_settings.custom_cfg is not None: model = model.sample_settings.custom_cfg.patch_model(model) # clone params from model diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index 11ca50a..ddadcc1 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -251,6 +251,11 @@ class ADKeyframeGroup: return cloned +class LoraHookMode: + MIN_VRAM = "min_vram" + MAX_SPEED = "max_speed" + + class LoraHook: def __init__(self, lora_name: str): self.lora_name = lora_name @@ -265,8 +270,14 @@ class LoraHookGroup: Stores LoRA hooks to apply for conditioning ''' def __init__(self): - self.hooks = [] + self.hooks: list[LoraHook] = [] + def names(self): + names = [] + for hook in self.hooks: + names.append(hook.lora_name) + return ",".join(names) + def add(self, hook: str): if hook not in self.hooks: self.hooks.append(hook) From b2b4f1ce5a2d1d22204155a00c9b489e7d91006d Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 27 Mar 2024 13:47:58 -0500 Subject: [PATCH 08/30] Added Register Model as LoRA Hook functionality --- animatediff/model_injection.py | 92 +++++++++++++++++++++++- animatediff/nodes.py | 14 +++- animatediff/nodes_conditioning.py | 113 +++++++++++++++++++++++++++++- animatediff/utils_motion.py | 16 +++++ 4 files changed, 229 insertions(+), 6 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 4e07d76..816dba9 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -7,6 +7,7 @@ import torch.nn.functional as F import torch import uuid +import comfy.lora import comfy.model_management import comfy.utils from comfy.model_patcher import ModelPatcher @@ -50,6 +51,8 @@ class ModelPatcherAndInjector(ModelPatcher): self.cached_hooked_patches = {} # binds LoraHookGroup to pre-calculated weights (speed optimization) self.current_lora_hooks = None self.lora_hook_mode = LoraHookMode.MAX_SPEED + # replace hook stuff + self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys # injection stuff self.currently_injected = False self.motion_injection_params: InjectionParams = InjectionParams() @@ -68,6 +71,11 @@ class ModelPatcherAndInjector(ModelPatcher): cloned.cached_hooked_patches[group] = {} for k in self.cached_hooked_patches[group]: cloned.cached_hooked_patches[group][k] = self.cached_hooked_patches[group][k] + # copy replace lora hooks + for hook in self.hooked_replace_patches: + cloned.hooked_replace_patches[hook] = {} + for k in self.hooked_replace_patches[hook]: + cloned.hooked_replace_patches[hook][k] = self.hooked_replace_patches[hook][k] cloned.hooked_backup = self.hooked_backup cloned.current_lora_hooks = self.current_lora_hooks cloned.currently_injected = self.currently_injected @@ -102,7 +110,7 @@ class ModelPatcherAndInjector(ModelPatcher): current_patches.append((strength_patch, patches[key], strength_model)) current_hooked_patches[key] = current_patches self.hooked_patches[lora_hook] = current_hooked_patches - # since should care about these patches too to determine if same model, roll patches_uuid + # since should care about these patches too to determine if same model, reroll patches_uuid self.patches_uuid = uuid.uuid4() return list(p) @@ -121,6 +129,27 @@ class ModelPatcherAndInjector(ModelPatcher): combined_patches[key] = current_patches return combined_patches + def add_hooked_replace_patches(self, lora_hook: LoraHook, patches: dict): + self.hooked_replace_patches.setdefault(lora_hook, {}) + p = set() + for key in patches: + if key in self.model_keys: + p.add(key) + self.hooked_replace_patches[lora_hook][key] = patches[key] + # since should care about these patches too to determine if same model, reroll patches_uuid + self.patches_uuid = uuid.uuid4() + return list(p) + + def get_hooked_replace_patches(self, lora_hooks: LoraHookGroup): + # return first hook found in hooked_replace_patches + patches = {} + if lora_hooks is not None: + for hook in lora_hooks.hooks: + if hook in self.hooked_replace_patches: + patches = self.hooked_replace_patches[hook] + break + return patches + def model_patches_to(self, device): super().model_patches_to(device) if self.motion_models is not None: @@ -208,9 +237,12 @@ class ModelPatcherAndInjector(ModelPatcher): else: # get combined patches of relevant lora_hooks relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) + replace_patches = self.get_hooked_replace_patches(lora_hooks=lora_hooks) + if len(replace_patches) > 0: + self.patch_hooked_replace_weight_to_device(lora_hooks=lora_hooks, model_sd=model_sd, replace_patches=replace_patches) for key in relevant_patches: if key not in model_sd: - logger.warning(f"LoraHook hook could not patch. key doesn't exist in model: {key}") + logger.warning(f"LoraHook could not patch. key doesn't exist in model: {key}") continue self.patch_hooked_weight_to_device(lora_hooks=lora_hooks, combined_patches=relevant_patches, key=key, device_to=device_to) self.current_lora_hooks = lora_hooks @@ -268,6 +300,29 @@ class ModelPatcherAndInjector(ModelPatcher): else: comfy.utils.set_attr_param(self.model, key, out_weight) + def patch_hooked_replace_weight_to_device(self, lora_hooks: LoraHookGroup, model_sd: dict, replace_patches: dict, device_to=None): + # first handle replace_patches + for key in replace_patches: + if key not in model_sd: + logger.warning(f"LoraHook could not replace patch. key doesn't exist in model: {key}") + continue + weight: Tensor = comfy.utils.get_attr(self.model, key) + inplace_update = self.weight_inplace_update + target_device = self.offload_device + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + target_device = self.current_device + + if key not in self.hooked_backup: + self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) + out_weight = replace_patches[key].to(self.load_device) + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + self.cached_hooked_patches.setdefault(lora_hooks, {}) + self.cached_hooked_patches[lora_hooks][key] = out_weight + if inplace_update: + comfy.utils.copy_to_param(self.model, key, out_weight) + else: + comfy.utils.set_attr_param(self.model, key, out_weight) + def unpatch_hooked(self, device_to=None, unpatch_weights=True) -> None: if not unpatch_weights: return @@ -323,6 +378,9 @@ class CLIPWithHooks(CLIP): 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 add_hooked_replace_patches(self, lora_hook: LoraHook, patches): + return self.patcher.add_hooked_replace_patches(lora_hook=lora_hook, patches=patches) + def load_model(self, *args, **kwargs): self.patcher.unpatch_hooked() returned = super().load_model(*args, **kwargs) @@ -359,12 +417,42 @@ def load_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInject k = set(k) k1 = set(k1) for x in loaded: + #if x.startswith() if (x not in k) and (x not in k1): logger.warning(f"NOT LOADED {x}") return (new_modelpatcher, new_clip) +def load_model_as_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInjector], clip: CLIP, model_loaded: ModelPatcher, clip_loaded: CLIP, lora_hook: LoraHook): + if model is not None and model_loaded is not None: + new_modelpatcher = ModelPatcherAndInjector.create_from(model) + k = new_modelpatcher.add_hooked_replace_patches(lora_hook=lora_hook, patches=model_loaded.model.state_dict()) + else: + k = () + new_modelpatcher = None + + if clip is not None and clip_loaded is not None: + new_clip = CLIPWithHooks(clip) + k1 = new_clip.add_hooked_replace_patches(lora_hook=lora_hook, patches=clip.cond_stage_model.state_dict()) + else: + k1 = () + new_clip = None + + k = set(k) + k1 = set(k1) + if model is not None and model_loaded is not None: + for key in model_loaded.model_keys: + if key not in k: + logger.warning(f"MODEL-AS-LORA NOT LOADED {key}") + if clip is not None and clip_loaded is not None: + for key in clip_loaded.patcher.model_keys: + if key not in k1: + logger.warning(f"CLIP-AS-LORA NOT LOADED {key}") + + return (new_modelpatcher, new_clip) + + class MotionModelPatcher(ModelPatcher): # Mostly here so that type hints work in IDEs def __init__(self, *args, **kwargs): diff --git a/animatediff/nodes.py b/animatediff/nodes.py index be7b9fc..7f431b2 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -6,7 +6,9 @@ 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 MaskableLoraLoader, MaskableLoraLoaderModelOnly, SetModelLoraHook, SetClipLoraHook, CombineLoraHooks +from .nodes_conditioning import (MaskableLoraLoader, MaskableLoraLoaderModelOnly, MaskableSDModelLoader, MaskableSDModelLoaderModelOnly, + SetModelLoraHook, SetClipLoraHook, + CombineLoraHooks, CombineLoraHookFourOptional, CombineLoraHookEightOptional) from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGKeyframeNode) from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode) @@ -52,7 +54,11 @@ NODE_CLASS_MAPPINGS = { # Conditioning "ADE_RegisterLoraHook": MaskableLoraLoader, "ADE_RegisterLoraHookModelOnly": MaskableLoraLoaderModelOnly, + #"ADE_RegisterModelAsLoraHook": MaskableSDModelLoader, # CLIP replace does not work properly + "ADE_RegisterModelAsLoraHookModelOnly": MaskableSDModelLoaderModelOnly, "ADE_CombineLoraHooks": CombineLoraHooks, + "ADE_CombineLoraHooksFour": CombineLoraHookFourOptional, + "ADE_CombineLoraHooksEight": CombineLoraHookEightOptional, "ADE_AttachLoraHookToConditioning": SetModelLoraHook, "ADE_AttachLoraHookToCLIP": SetClipLoraHook, # Noise Layer Nodes @@ -129,7 +135,11 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Conditioning "ADE_RegisterLoraHook": "Register LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_RegisterLoraHookModelOnly": "Register LoRA Hook (Model Only) πŸŽ­πŸ…πŸ…“", - "ADE_CombineLoraHooks": "Combine LoRA Hooks πŸŽ­πŸ…πŸ…“", + #"ADE_RegisterModelAsLoraHook": "Register Model as LoRA Hook+ πŸŽ­πŸ…πŸ…“", # CLIP replace does not work properly + "ADE_RegisterModelAsLoraHookModelOnly": "Register Model as LoRA Hook πŸŽ­πŸ…πŸ…“", + "ADE_CombineLoraHooks": "Combine LoRA Hooks [2] πŸŽ­πŸ…πŸ…“", + "ADE_CombineLoraHooksFour": "Combine LoRA Hooks [4] πŸŽ­πŸ…πŸ…“", + "ADE_CombineLoraHooksEight": "Combine LoRA Hooks [8] πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 32fb82a..31d8c43 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -4,10 +4,11 @@ from typing import Union from comfy.model_patcher import ModelPatcher from comfy.sd import CLIP +import comfy.sd import comfy.utils from .utils_motion import LoraHook, LoraHookGroup -from .model_injection import ModelPatcherAndInjector, CLIPWithHooks, load_hooked_lora_for_models +from .model_injection import ModelPatcherAndInjector, CLIPWithHooks, load_hooked_lora_for_models, load_model_as_hooked_lora_for_models # based on ComfyUI's nodes.py LoraLoader class MaskableLoraLoader: @@ -77,6 +78,55 @@ class MaskableLoraLoaderModelOnly(MaskableLoraLoader): return (model_lora, lora_hook) +class MaskableSDModelLoader: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "clip": ("CLIP",), + "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ), + } + } + + RETURN_TYPES = ("MODEL", "CLIP", "LORA_HOOK") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "load_model_as_lora" + + def load_model_as_lora(self, model: ModelPatcher, clip: CLIP, ckpt_name: str): + ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) + out = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) + model_loaded = out[0] + clip_loaded = out[1] + + lora_hook = LoraHook(lora_name=ckpt_name) + lora_hook_group = LoraHookGroup() + lora_hook_group.add(lora_hook) + model_lora, clip_lora = load_model_as_hooked_lora_for_models(model=model, clip=clip, + model_loaded=model_loaded, clip_loaded=clip_loaded, + lora_hook=lora_hook) + return (model_lora, clip_lora, lora_hook_group) + + +class MaskableSDModelLoaderModelOnly(MaskableSDModelLoader): + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("MODEL",), + "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ), + } + } + + RETURN_TYPES = ("MODEL", "LORA_HOOK") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "load_model_as_lora_model_only" + + def load_model_as_lora_model_only(self, model: ModelPatcher, ckpt_name: str): + model_lora, clip_lora, lora_hook = self.load_model_as_lora(model=model, clip=None, ckpt_name=ckpt_name) + return (model_lora, lora_hook) + + class SetModelLoraHook: @classmethod def INPUT_TYPES(s): @@ -132,8 +182,67 @@ class CombineLoraHooks: } RETURN_TYPES = ("LORA_HOOK",) - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine hooks" FUNCTION = "combine_lora_hooks" def combine_lora_hooks(self, lora_hook_A: LoraHookGroup, lora_hook_B: LoraHookGroup): return (lora_hook_A.clone_and_combine(lora_hook_B),) + + +class CombineLoraHookFourOptional: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + + }, + "optional": { + "lora_hook_A": ("LORA_HOOK",), + "lora_hook_B": ("LORA_HOOK",), + "lora_hook_C": ("LORA_HOOK",), + "lora_hook_D": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("LORA_HOOK",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine hooks" + FUNCTION = "combine_lora_hooks" + + def combine_lora_hooks(self, + lora_hook_A: LoraHookGroup=None, lora_hook_B: LoraHookGroup=None, + lora_hook_C: LoraHookGroup=None, lora_hook_D: LoraHookGroup=None,): + candidates = [lora_hook_A, lora_hook_B, lora_hook_C, lora_hook_D] + return (LoraHookGroup.combine_all_lora_hooks(candidates),) + + +class CombineLoraHookEightOptional: + @classmethod + def INPUT_TYPES(s): + return { + "optional": { + "lora_hook_A": ("LORA_HOOK",), + "lora_hook_B": ("LORA_HOOK",), + "lora_hook_C": ("LORA_HOOK",), + "lora_hook_D": ("LORA_HOOK",), + "lora_hook_E": ("LORA_HOOK",), + "lora_hook_F": ("LORA_HOOK",), + "lora_hook_G": ("LORA_HOOK",), + "lora_hook_H": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("LORA_HOOK",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine hooks" + FUNCTION = "combine_lora_hooks" + + def combine_lora_hooks(self, + lora_hook_A: LoraHookGroup=None, lora_hook_B: LoraHookGroup=None, + lora_hook_C: LoraHookGroup=None, lora_hook_D: LoraHookGroup=None, + lora_hook_E: LoraHookGroup=None, lora_hook_F: LoraHookGroup=None, + lora_hook_G: LoraHookGroup=None, lora_hook_H: LoraHookGroup=None): + candidates = [lora_hook_A, lora_hook_B, lora_hook_C, lora_hook_D, + lora_hook_E, lora_hook_F, lora_hook_G, lora_hook_H] + return (LoraHookGroup.combine_all_lora_hooks(candidates),) + +# NOTE: if at some point I add more Javascript stuff to this repo, there should be a combine node +# that dynamically increases the hooks available to plug in on the node diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index ddadcc1..1ed0cd6 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -297,6 +297,22 @@ class LoraHookGroup: cloned.add(hook) return cloned + @staticmethod + def combine_all_lora_hooks(lora_hooks_list: list['LoraHookGroup'], require_count=2) -> 'LoraHookGroup': + actual: list[LoraHookGroup] = [] + for group in lora_hooks_list: + if group is not None: + actual.append(group) + if len(actual) < require_count: + raise Exception(f"Need at least {require_count} LoRA Hooks to combine, but only had {len(actual)}.") + final_hook: LoraHookGroup = None + for hook in actual: + if final_hook is None: + final_hook = hook.clone() + else: + final_hook = final_hook.clone_and_combine(hook) + return final_hook + class DummyNNModule(nn.Module): class DoNothingWhenCalled: From d45bc705a4b34641d96a0e2bc5a732509a6664ac Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 28 Mar 2024 01:30:55 -0500 Subject: [PATCH 09/30] Improved CLIP hooked LoRA patching, fixed eight-way Combine LoRA Hooks node not working --- animatediff/model_injection.py | 50 +++++++++++++++++++++++-------- animatediff/nodes.py | 24 +++++++-------- animatediff/nodes_conditioning.py | 3 +- 3 files changed, 51 insertions(+), 26 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 816dba9..86b13fe 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -1,5 +1,5 @@ import copy -from typing import Union +from typing import Union, Callable from einops import rearrange from torch import Tensor @@ -381,17 +381,43 @@ class CLIPWithHooks(CLIP): def add_hooked_replace_patches(self, lora_hook: LoraHook, patches): return self.patcher.add_hooked_replace_patches(lora_hook=lora_hook, patches=patches) - 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 load_model(self, *args, **kwargs): + # #comfy.model_management.cleanup_models() + # #return super().load_model(*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 - + # to work properly, need to hack (and then unhack) cond_stage_model's encode_token_weights function + def encode_token_weights_factory(orig_encode_token_weights: Callable, clip: CLIPWithHooks): + def encode_token_weights_wrapper_hooked(*args, **kwargs): + try: + # yeah, not sure why, but first patching is screwed up if things change... BUT only the first time. + # so we apply it twice to avoid it... this will be future me's problem. + # the first result will still end up slightly different than intended, but it's very close. + if clip.desired_hooks is not None: + temp_desired_hooks = clip.desired_hooks.clone() if clip.desired_hooks is not None else None + clip.patcher.clear_cached_hooked_weights() + clip.patcher.unpatch_hooked() + clip.patcher.apply_lora_hooks(temp_desired_hooks) + orig_encode_token_weights(*args, **kwargs) + clip.patcher.clear_cached_hooked_weights() + clip.patcher.apply_lora_hooks(clip.desired_hooks) + return orig_encode_token_weights(*args, **kwargs) + finally: + clip.patcher.unpatch_hooked() + clip.patcher.clear_cached_hooked_weights() + return encode_token_weights_wrapper_hooked + try: + orig_encode_token_weights = self.cond_stage_model.encode_token_weights + self.cond_stage_model.encode_token_weights = encode_token_weights_factory(orig_encode_token_weights, self) + return super().encode_from_tokens(*args, **kwargs) + finally: + self.cond_stage_model.encode_token_weights = orig_encode_token_weights + 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 = {} @@ -400,7 +426,7 @@ def load_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInject if clip is not None: key_map = comfy.lora.model_lora_keys_clip(clip.cond_stage_model, key_map) - loaded = comfy.lora.load_lora(lora, key_map) + loaded: dict[str] = comfy.lora.load_lora(lora, key_map) if model is not None: new_modelpatcher = ModelPatcherAndInjector.create_from(model) k = new_modelpatcher.add_hooked_patches(lora_hook=lora_hook, patches=loaded, strength_patch=strength_model) @@ -417,10 +443,8 @@ def load_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInject k = set(k) k1 = set(k1) for x in loaded: - #if x.startswith() if (x not in k) and (x not in k1): logger.warning(f"NOT LOADED {x}") - return (new_modelpatcher, new_clip) diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 7f431b2..6fd87f8 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -54,7 +54,7 @@ NODE_CLASS_MAPPINGS = { # Conditioning "ADE_RegisterLoraHook": MaskableLoraLoader, "ADE_RegisterLoraHookModelOnly": MaskableLoraLoaderModelOnly, - #"ADE_RegisterModelAsLoraHook": MaskableSDModelLoader, # CLIP replace does not work properly + #"ADE_RegisterModelAsLoraHook": MaskableSDModelLoader, # CLIP does not work properly on first run "ADE_RegisterModelAsLoraHookModelOnly": MaskableSDModelLoaderModelOnly, "ADE_CombineLoraHooks": CombineLoraHooks, "ADE_CombineLoraHooksFour": CombineLoraHookFourOptional, @@ -91,10 +91,6 @@ NODE_CLASS_MAPPINGS = { # Gen1 Nodes "ADE_AnimateDiffLoaderGen1": AnimateDiffLoaderGen1, "ADE_AnimateDiffLoaderWithContext": LegacyAnimateDiffLoaderWithContext, - "ADE_AnimateDiffModelSettings_Release": AnimateDiffModelSettings, - "ADE_AnimateDiffModelSettingsSimple": AnimateDiffModelSettingsSimple, - "ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced, - "ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": AnimateDiffModelSettingsAdvancedAttnStrengths, # Gen2 Nodes "ADE_UseEvolvedSampling": UseEvolvedSamplingNode, "ADE_ApplyAnimateDiffModelSimple": ApplyAnimateDiffModelBasicNode, @@ -109,6 +105,10 @@ NODE_CLASS_MAPPINGS = { "AnimateDiffLoaderV1": AnimateDiffLoader_Deprecated, "ADE_AnimateDiffLoaderV1Advanced": AnimateDiffLoaderAdvanced_Deprecated, "ADE_AnimateDiffCombine": AnimateDiffCombine_Deprecated, + "ADE_AnimateDiffModelSettings_Release": AnimateDiffModelSettings, + "ADE_AnimateDiffModelSettingsSimple": AnimateDiffModelSettingsSimple, + "ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced, + "ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": AnimateDiffModelSettingsAdvancedAttnStrengths, } NODE_DISPLAY_NAME_MAPPINGS = { # Unencapsulated @@ -135,13 +135,13 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Conditioning "ADE_RegisterLoraHook": "Register LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_RegisterLoraHookModelOnly": "Register LoRA Hook (Model Only) πŸŽ­πŸ…πŸ…“", - #"ADE_RegisterModelAsLoraHook": "Register Model as LoRA Hook+ πŸŽ­πŸ…πŸ…“", # CLIP replace does not work properly - "ADE_RegisterModelAsLoraHookModelOnly": "Register Model as LoRA Hook πŸŽ­πŸ…πŸ…“", + #"ADE_RegisterModelAsLoraHook": "πŸ”¬Register Model as LoRA Hook πŸŽ­πŸ…πŸ…“", # CLIP does not work properly on first run + "ADE_RegisterModelAsLoraHookModelOnly": "Register Model as LoRA Hook (MO) πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooks": "Combine LoRA Hooks [2] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksFour": "Combine LoRA Hooks [4] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksEight": "Combine LoRA Hooks [8] πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", - "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", + "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA HookπŸ”¬ πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", @@ -172,10 +172,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Gen1 Nodes "ADE_AnimateDiffLoaderGen1": "AnimateDiff Loader πŸŽ­πŸ…πŸ…“β‘ ", "ADE_AnimateDiffLoaderWithContext": "AnimateDiff Loader [Legacy] πŸŽ­πŸ…πŸ…“β‘ ", - "ADE_AnimateDiffModelSettings_Release": "🚫[DEPR] Motion Model Settings πŸŽ­πŸ…πŸ…“β‘ ", - "ADE_AnimateDiffModelSettingsSimple": "🚫[DEPR] Motion Model Settings (Simple) πŸŽ­πŸ…πŸ…“β‘ ", - "ADE_AnimateDiffModelSettings": "🚫[DEPR] Motion Model Settings (Advanced) πŸŽ­πŸ…πŸ…“β‘ ", - "ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": "🚫[DEPR] Motion Model Settings (Adv. Attn) πŸŽ­πŸ…πŸ…“β‘ ", # Gen2 Nodes "ADE_UseEvolvedSampling": "Use Evolved Sampling πŸŽ­πŸ…πŸ…“β‘‘", "ADE_ApplyAnimateDiffModelSimple": "Apply AnimateDiff Model πŸŽ­πŸ…πŸ…“β‘‘", @@ -190,4 +186,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "AnimateDiffLoaderV1": "🚫AnimateDiff Loader [DEPRECATED] πŸŽ­πŸ…πŸ…“", "ADE_AnimateDiffLoaderV1Advanced": "🚫AnimateDiff Loader (Advanced) [DEPRECATED] πŸŽ­πŸ…πŸ…“", "ADE_AnimateDiffCombine": "🚫AnimateDiff Combine [DEPRECATED, Use Video Combine (VHS) Instead!] πŸŽ­πŸ…πŸ…“", + "ADE_AnimateDiffModelSettings_Release": "🚫[DEPR] Motion Model Settings πŸŽ­πŸ…πŸ…“β‘ ", + "ADE_AnimateDiffModelSettingsSimple": "🚫[DEPR] Motion Model Settings (Simple) πŸŽ­πŸ…πŸ…“β‘ ", + "ADE_AnimateDiffModelSettings": "🚫[DEPR] Motion Model Settings (Advanced) πŸŽ­πŸ…πŸ…“β‘ ", + "ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": "🚫[DEPR] Motion Model Settings (Adv. Attn) πŸŽ­πŸ…πŸ…“β‘ ", } diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 31d8c43..d759fb7 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -194,7 +194,6 @@ class CombineLoraHookFourOptional: def INPUT_TYPES(s): return { "required": { - }, "optional": { "lora_hook_A": ("LORA_HOOK",), @@ -219,6 +218,8 @@ class CombineLoraHookEightOptional: @classmethod def INPUT_TYPES(s): return { + "required": { + }, "optional": { "lora_hook_A": ("LORA_HOOK",), "lora_hook_B": ("LORA_HOOK",), From 1e8fba93844a7426cb28ff281ddffdb498785650 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 29 Mar 2024 06:22:30 -0500 Subject: [PATCH 10/30] Made CLIP LoRA hooks work as intended, proper way to force ModelPatcherAndInjector to cause model reload when hooks present --- animatediff/model_injection.py | 248 +++++++++++++++++++++++++++------ animatediff/nodes.py | 4 +- animatediff/sampling.py | 2 - animatediff/utils_motion.py | 2 + 4 files changed, 213 insertions(+), 43 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 86b13fe..affb13c 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -93,6 +93,23 @@ class ModelPatcherAndInjector(ModelPatcher): else: return ModelPatcherAndInjector(model) + def clone_has_same_weights(self, clone: 'ModelPatcherCLIPHooks'): + returned = super().clone_has_same_weights(clone) + if not returned: + return returned + # currently, hook patches require that model gets loaded when sampled, so always say is not a clone if hooks present + if len(self.hooked_patches) > 0 or len(self.hooked_replace_patches) > 0: + return False + if type(self) != type(clone): + return False + if self.current_lora_hooks != clone.current_lora_hooks: + return False + if self.hooked_patches.keys() != clone.hooked_patches.keys(): + return False + if self.hooked_replace_patches.keys() != clone.hooked_replace_patches.keys(): + return False + return returned + def set_lora_hook_mode(self, lora_hook_mode: str): self.lora_hook_mode = lora_hook_mode @@ -208,7 +225,6 @@ class ModelPatcherAndInjector(ModelPatcher): # unpatch hooks, if needed self.unpatch_hooked(device_to=device_to) # finally, patch hooks - # TODO: handle lowvram self.patch_hooked(lora_hooks=lora_hooks, device_to=device_to) def patch_hooked(self, lora_hooks: LoraHookGroup, device_to=None, patch_weights=True) -> None: @@ -235,6 +251,7 @@ class ModelPatcherAndInjector(ModelPatcher): logger.warning(f"Cached LoraHook hook could not patch. key doesn't exist in model: {key}") self.patch_cached_hooked_weight(cached_weights=cached_weights, key=key) else: + # TODO: handle lowvram # get combined patches of relevant lora_hooks relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) replace_patches = self.get_hooked_replace_patches(lora_hooks=lora_hooks) @@ -255,6 +272,7 @@ class ModelPatcherAndInjector(ModelPatcher): pass def patch_cached_hooked_weight(self, cached_weights: dict, key: str): + # TODO: handle lowvram inplace_update = self.weight_inplace_update target_device = self.offload_device if self.lora_hook_mode == LoraHookMode.MAX_SPEED: @@ -359,14 +377,13 @@ class ModelPatcherAndInjector(ModelPatcher): class CLIPWithHooks(CLIP): def __init__(self, clip: Union[CLIP, 'CLIPWithHooks']): super().__init__(no_init=True) - self.patcher = ModelPatcherAndInjector.create_from(clip.patcher) - self.patcher.set_lora_hook_mode(lora_hook_mode=LoraHookMode.MIN_VRAM) + self.patcher = ModelPatcherCLIPHooks.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 + self.set_desired_hooks(clip.desired_hooks) def clone(self): cloned = CLIPWithHooks(clip=self) @@ -374,6 +391,7 @@ class CLIPWithHooks(CLIP): def set_desired_hooks(self, lora_hooks: LoraHookGroup): self.desired_hooks = lora_hooks + self.patcher.set_desired_hooks(lora_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) @@ -381,43 +399,195 @@ class CLIPWithHooks(CLIP): def add_hooked_replace_patches(self, lora_hook: LoraHook, patches): return self.patcher.add_hooked_replace_patches(lora_hook=lora_hook, patches=patches) - # def load_model(self, *args, **kwargs): - # #comfy.model_management.cleanup_models() - # #return super().load_model(*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, tokens, return_pooled=False): + # comfy.model_management.cleanup_models() + # return super().encode_from_tokens(tokens, return_pooled) - def encode_from_tokens(self, *args, **kwargs): - # to work properly, need to hack (and then unhack) cond_stage_model's encode_token_weights function - def encode_token_weights_factory(orig_encode_token_weights: Callable, clip: CLIPWithHooks): - def encode_token_weights_wrapper_hooked(*args, **kwargs): - try: - # yeah, not sure why, but first patching is screwed up if things change... BUT only the first time. - # so we apply it twice to avoid it... this will be future me's problem. - # the first result will still end up slightly different than intended, but it's very close. - if clip.desired_hooks is not None: - temp_desired_hooks = clip.desired_hooks.clone() if clip.desired_hooks is not None else None - clip.patcher.clear_cached_hooked_weights() - clip.patcher.unpatch_hooked() - clip.patcher.apply_lora_hooks(temp_desired_hooks) - orig_encode_token_weights(*args, **kwargs) - clip.patcher.clear_cached_hooked_weights() - clip.patcher.apply_lora_hooks(clip.desired_hooks) - return orig_encode_token_weights(*args, **kwargs) - finally: - clip.patcher.unpatch_hooked() - clip.patcher.clear_cached_hooked_weights() - return encode_token_weights_wrapper_hooked - try: - orig_encode_token_weights = self.cond_stage_model.encode_token_weights - self.cond_stage_model.encode_token_weights = encode_token_weights_factory(orig_encode_token_weights, self) - return super().encode_from_tokens(*args, **kwargs) - finally: - self.cond_stage_model.encode_token_weights = orig_encode_token_weights + +class ModelPatcherCLIPHooks(ModelPatcher): + def __init__(self, m: ModelPatcher): + # replicate ModelPatcher.clone() to initialize + super().__init__(m.model, m.load_device, m.offload_device, m.size, m.current_device, weight_inplace_update=m.weight_inplace_update) + self.patches = {} + for k in m.patches: + self.patches[k] = m.patches[k][:] + if hasattr(m, "patches_uuid"): + self.patches_uuid = m.patches_uuid + + self.object_patches = m.object_patches.copy() + self.model_options = copy.deepcopy(m.model_options) + self.model_keys = m.model_keys + if hasattr(m, "backup"): + self.backup = m.backup + if hasattr(m, "object_patches_backup"): + self.object_patches_backup = m.object_patches_backup + # lora hook stuff + self.hooked_patches = {} # binds LoraHook to specific keys + self.patches_backup = {} + self.hooked_backup = {} + + self.current_lora_hooks = None + self.desired_lora_hooks = None + self.lora_hook_mode = LoraHookMode.MAX_SPEED + # replace hook stuff + self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys + + def clone(self): + cloned = ModelPatcherCLIPHooks(self) + # copy lora hooks + for hook in self.hooked_patches: + cloned.hooked_patches[hook] = {} + for k in self.hooked_patches[hook]: + cloned.hooked_patches[hook][k] = self.hooked_patches[hook][k][:] + # copy replace lora hooks + for hook in self.hooked_replace_patches: + cloned.hooked_replace_patches[hook] = {} + for k in self.hooked_replace_patches[hook]: + cloned.hooked_replace_patches[hook][k] = self.hooked_replace_patches[hook][k] + cloned.patches_backup = self.patches_backup + cloned.hooked_backup = self.hooked_backup + cloned.current_lora_hooks = self.current_lora_hooks + cloned.desired_lora_hooks = self.desired_lora_hooks + cloned.lora_hook_mode = self.lora_hook_mode + return cloned + + @classmethod + def create_from(cls, model: Union[ModelPatcher, 'ModelPatcherCLIPHooks']): + if isinstance(model, ModelPatcherCLIPHooks): + return model.clone() + return ModelPatcherCLIPHooks(model) + + def clone_has_same_weights(self, clone: 'ModelPatcherCLIPHooks'): + returned = super().clone_has_same_weights(clone) + if not returned: + return returned + if type(self) != type(clone): + return False + if self.desired_lora_hooks != clone.desired_lora_hooks: + return False + if self.current_lora_hooks != clone.current_lora_hooks: + return False + if self.hooked_patches.keys() != clone.hooked_patches.keys(): + return False + if self.hooked_replace_patches.keys() != clone.hooked_replace_patches.keys(): + return False + return returned + + def set_desired_hooks(self, lora_hooks: LoraHookGroup): + self.desired_lora_hooks = lora_hooks + + def add_hooked_patches(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): + ''' + Based on add_patches, but for hooked weights. + ''' + # TODO: make this work with timestep scheduling + current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {}) + p = set() + for key in patches: + if key in self.model_keys: + p.add(key) + current_patches: list[tuple] = current_hooked_patches.get(key, []) + current_patches.append((strength_patch, patches[key], strength_model)) + current_hooked_patches[key] = current_patches + self.hooked_patches[lora_hook] = current_hooked_patches + # since should care about these patches too to determine if same model, reroll patches_uuid + self.patches_uuid = uuid.uuid4() + return list(p) + + def get_combined_hooked_patches(self, lora_hooks: LoraHookGroup): + ''' + Returns patches for selected lora_hooks. + ''' + # combined_patches will contain weights of all relevant lora_hooks, per key + combined_patches = {} + if lora_hooks is not None: + for hook in lora_hooks.hooks: + hook_patches: dict = self.hooked_patches.get(hook, {}) + for key in hook_patches.keys(): + current_patches: list[tuple] = combined_patches.get(key, []) + current_patches.extend(hook_patches[key]) + combined_patches[key] = current_patches + return combined_patches + + def add_hooked_replace_patches(self, lora_hook: LoraHook, patches: dict): + self.hooked_replace_patches.setdefault(lora_hook, {}) + p = set() + for key in patches: + if key in self.model_keys: + p.add(key) + self.hooked_replace_patches[lora_hook][key] = patches[key] + # since should care about these patches too to determine if same model, reroll patches_uuid + self.patches_uuid = uuid.uuid4() + return list(p) + + def get_hooked_replace_patches(self, lora_hooks: LoraHookGroup): + # return first hook found in hooked_replace_patches + patches = {} + if lora_hooks is not None: + for hook in lora_hooks.hooks: + if hook in self.hooked_replace_patches: + patches = self.hooked_replace_patches[hook] + break + return patches + + def patch_hooked_replace_weight_to_device(self, model_sd: dict, replace_patches: dict): + # first handle replace_patches + for key in replace_patches: + if key not in model_sd: + logger.warning(f"CLIP LoraHook could not replace patch. key doesn't exist in model: {key}") + continue + weight: Tensor = comfy.utils.get_attr(self.model, key) + inplace_update = self.weight_inplace_update + target_device = self.current_device + if key not in self.hooked_backup: + self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) + out_weight = replace_patches[key].to(target_device) + if inplace_update: + comfy.utils.copy_to_param(self.model, key, out_weight) + else: + comfy.utils.set_attr_param(self.model, key, out_weight) + + def patch_model(self, device_to=None, patch_weights=True, *args, **kwargs): + if self.desired_lora_hooks is not None: + self.patches_backup = self.patches.copy() + # first, handle replace patches # TODO: make work properly for CLIP + replace_patches = self.get_hooked_replace_patches(lora_hooks=self.desired_lora_hooks) + if len(replace_patches) > 0: + model_sd = self.model_state_dict() + self.patch_hooked_replace_weight_to_device(model_sd=model_sd, replace_patches=replace_patches) + # then, handle usual patches + relevant_patches = self.get_combined_hooked_patches(lora_hooks=self.desired_lora_hooks) + for key in relevant_patches: + self.patches.setdefault(key, []) + self.patches[key].extend(relevant_patches[key]) + self.current_lora_hooks = self.desired_lora_hooks + return super().patch_model(device_to, patch_weights, *args, **kwargs) + + def unpatch_model(self, device_to=None, unpatch_weights=True, *args, **kwargs): + try: + return super().unpatch_model(device_to, unpatch_weights, *args, **kwargs) + finally: + self.patches = self.patches_backup.copy() + self.patches_backup.clear() + # handle replace patches + keys = list(self.hooked_backup.keys()) + if self.weight_inplace_update: + for k in keys: + if device_to is None: + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) + else: + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k]) + else: + for k in keys: + if device_to is None: + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) + else: + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k]) + # clear hooked_backup + self.hooked_backup.clear() + self.current_lora_hooks = None + 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 = {} diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 6fd87f8..4ef4bc2 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -135,13 +135,13 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Conditioning "ADE_RegisterLoraHook": "Register LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_RegisterLoraHookModelOnly": "Register LoRA Hook (Model Only) πŸŽ­πŸ…πŸ…“", - #"ADE_RegisterModelAsLoraHook": "πŸ”¬Register Model as LoRA Hook πŸŽ­πŸ…πŸ…“", # CLIP does not work properly on first run + #"ADE_RegisterModelAsLoraHook": "Register Model as LoRA HookπŸ”¬ πŸŽ­πŸ…πŸ…“", # CLIP does not work properly on first run "ADE_RegisterModelAsLoraHookModelOnly": "Register Model as LoRA Hook (MO) πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooks": "Combine LoRA Hooks [2] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksFour": "Combine LoRA Hooks [4] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksEight": "Combine LoRA Hooks [8] πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", - "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA HookπŸ”¬ πŸŽ­πŸ…πŸ…“", + "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/sampling.py b/animatediff/sampling.py index e202179..f483ed1 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -277,8 +277,6 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) -> cached_noise = None function_injections = FunctionInjectionHolder() try: - if len(model.hooked_patches) > 0: - model_management.cleanup_models() if model.sample_settings.custom_cfg is not None: model = model.sample_settings.custom_cfg.patch_model(model) # clone params from model diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index 1ed0cd6..5b8f1f6 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -253,7 +253,9 @@ class ADKeyframeGroup: class LoraHookMode: MIN_VRAM = "min_vram" + MIN_VRAM_LOWVRAM = "min_vram_lowvram" MAX_SPEED = "max_speed" + MAX_SPEED_LOWVRAM = "max_speed_lowvram" class LoraHook: From 430c49ef04ba2ca2cf5a8adda6cec176f7cd7f03 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 3 Apr 2024 04:09:24 -0500 Subject: [PATCH 11/30] Use new calc_cond_batch function while maintaining backwards compatibility --- animatediff/sampling.py | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index f483ed1..b88e432 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -409,7 +409,7 @@ def evolved_sampling_function(model, x, timestep, uncond, cond, cond_scale, mode model_options["transformer_options"]["ad_params"] = ADGS.create_exposed_params() if not ADGS.is_using_sliding_context(): - cond_pred, uncond_pred = calc_cond_uncond_batch_wrapper(model, cond, uncond_, x, timestep, model_options) + cond_pred, uncond_pred = calc_cond_uncond_batch_wrapper(model, [cond, uncond_], x, timestep, model_options) else: cond_pred, uncond_pred = sliding_calc_cond_uncond_batch(model, cond, uncond_, x, timestep, model_options) @@ -428,6 +428,7 @@ def evolved_sampling_function(model, x, timestep, uncond, cond, cond_scale, mode return cfg_result +# TODO: properly carry over changes from calc_cond_batch # sliding_calc_cond_uncond_batch inspired by ashen's initial hack for 16-frame sliding context: # https://github.com/comfyanonymous/ComfyUI/compare/master...ashen-sensored:ComfyUI:master def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, model_options): @@ -523,7 +524,7 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, sub_cond = get_resized_cond(cond, full_idxs, len(ctx_idxs)) if cond is not None else None sub_uncond = get_resized_cond(uncond, full_idxs, len(ctx_idxs)) if uncond is not None else None - sub_cond_out, sub_uncond_out = calc_cond_uncond_batch_wrapper(model, sub_cond, sub_uncond, sub_x, sub_timestep, model_options) + sub_cond_out, sub_uncond_out = calc_cond_uncond_batch_wrapper(model, [sub_cond, sub_uncond], sub_x, sub_timestep, model_options) if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE: full_length = ADGS.params.full_length @@ -560,10 +561,10 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, return cond_final, uncond_final -def calc_cond_uncond_batch_wrapper(model, cond, uncond, x_in, timestep, model_options): +def calc_cond_uncond_batch_wrapper(model, conds: list[dict], x_in: Tensor, timestep, model_options): # check if conds or unconds contain lora_hook contains_lora_hooks = False - for cond_uncond in [cond, uncond]: + for cond_uncond in conds: if cond_uncond is None: continue for t in cond_uncond: @@ -573,9 +574,11 @@ def calc_cond_uncond_batch_wrapper(model, cond, uncond, x_in, timestep, model_op if contains_lora_hooks: break if contains_lora_hooks: - return calc_cond_uncond_batch_lora_hook(model, cond, uncond, x_in, timestep, model_options) - return comfy.samplers.calc_cond_uncond_batch(model, cond, uncond, x_in, timestep, model_options) - + return calc_cond_uncond_batch_lora_hook(model, conds[0], conds[1], x_in, timestep, model_options) + # keep for backwards compatibility, for now + if not hasattr(comfy.samplers, "calc_cond_batch"): + return comfy.samplers.calc_cond_uncond_batch(model, conds[0], conds[1], x_in, timestep, model_options) + return comfy.samplers.calc_cond_batch(model, conds, x_in, timestep, model_options) def calc_cond_uncond_batch_lora_hook(model, cond, uncond, x_in, timestep, model_options): out_cond = torch.zeros_like(x_in) From fa08393388f8d4146024db99f302afe5cf8b8cdb Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 23 Apr 2024 03:03:33 -0500 Subject: [PATCH 12/30] Fixed git merging mistake that happened at some point --- animatediff/sampling.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index bdc08fe..8f0fac9 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -642,7 +642,7 @@ def calc_cond_uncond_batch_lora_hook(model, cond, uncond, x_in, timestep, model_ to_batch_temp.reverse() to_batch = to_batch_temp[:1] - free_memory = model_management.get_free_memory(x_in.device) + free_memory = comfy.model_management.get_free_memory(x_in.device) for i in range(1, len(to_batch_temp) + 1): batch_amount = to_batch_temp[:len(to_batch_temp)//i] input_shape = [len(batch_amount) * first_shape[0]] + list(first_shape)[1:] From 89b33022b676ae8971d4ca6ab7f1113cd7e623dd Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 24 Apr 2024 07:12:30 -0500 Subject: [PATCH 13/30] Refactored sampling code to fully mirror calc_conds_batch, fixed only the top layer (last) Control object being checked if it is compatible with sliding context --- animatediff/sampling.py | 149 ++++++++++++++++++---------------------- 1 file changed, 68 insertions(+), 81 deletions(-) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index daedc3b..8160939 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -18,6 +18,7 @@ except ImportError: SAMPLE_FALLBACK = True import comfy.utils from comfy.controlnet import ControlBase +from comfy.model_base import BaseModel import comfy.ops from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows @@ -429,7 +430,7 @@ def evolved_sampling_function(model, x, timestep, uncond, cond, cond_scale, mode if not ADGS.is_using_sliding_context(): cond_pred, uncond_pred = calc_cond_uncond_batch_wrapper(model, [cond, uncond_], x, timestep, model_options) else: - cond_pred, uncond_pred = sliding_calc_cond_uncond_batch(model, cond, uncond_, x, timestep, model_options) + cond_pred, uncond_pred = sliding_calc_conds_batch(model, [cond, uncond_], x, timestep, model_options) if hasattr(comfy.samplers, "cfg_function"): try: @@ -460,34 +461,33 @@ def wrapped_cfg_sliding_calc_cond_batch_factory(orig_calc_cond_batch): def wrapped_cfg_sliding_calc_cond_batch(model, conds, x_in, timestep, model_options): # current call to calc_cond_batch should refer to sliding version try: - uncond = None current_calc_cond_batch = comfy.samplers.calc_cond_batch - # when inside sliding_calc_cond_uncond, should return to original calc_cond_batch + # when inside sliding_calc_conds_batch, should return to original calc_cond_batch comfy.samplers.calc_cond_batch = orig_calc_cond_batch - if len(conds) > 1: - uncond = conds[1] - result = sliding_calc_cond_uncond_batch(model, conds[0], uncond, x_in, timestep, model_options) - if uncond is None: - result = (result[0],) + result = sliding_calc_conds_batch(model, conds, x_in, timestep, model_options) return result finally: - del uncond # make sure calc_cond_batch will become wrapped again comfy.samplers.calc_cond_batch = current_calc_cond_batch return wrapped_cfg_sliding_calc_cond_batch -# sliding_calc_cond_uncond_batch inspired by ashen's initial hack for 16-frame sliding context: +# sliding_calc_conds_batch inspired by ashen's initial hack for 16-frame sliding context: # https://github.com/comfyanonymous/ComfyUI/compare/master...ashen-sensored:ComfyUI:master -def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, model_options): +def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options): def prepare_control_objects(control: ControlBase, full_idxs: list[int]): if control.previous_controlnet is not None: prepare_control_objects(control.previous_controlnet, full_idxs) + if not hasattr(control, "sub_idxs"): + raise ValueError(f"Control type {type(control).__name__} may not support required features for sliding context window; \ + use ControlNet nodes from Kosinkadink/ComfyUI-Advanced-ControlNet, or make sure ComfyUI-Advanced-ControlNet is updated.") control.sub_idxs = full_idxs control.full_latent_length = ADGS.params.full_length control.context_length = ADGS.params.context_options.context_length def get_resized_cond(cond_in, full_idxs: list[int], context_length: int) -> list: + if cond_in is None: + return None # reuse or resize cond items to match context requirements resized_cond = [] # cond object is a list containing a dict - outer list is irrelevant, so just loop through it @@ -545,14 +545,13 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, if ADGS.motion_models is not None: ADGS.motion_models.set_view_options(ADGS.params.context_options.view_options) + + # prepare final conds, out_counts, and biases + conds_final = [torch.zeros_like(x_in) for _ in conds] + counts_final = [torch.zeros((x_in.shape[0], 1, 1, 1), device=x_in.device) for _ in conds] + biases_final = [([0.0] * x_in.shape[0]) for _ in conds] - # prepare final cond, uncond, and out_count - cond_final = torch.zeros_like(x_in) - uncond_final = torch.zeros_like(x_in) - out_count_final = torch.zeros((x_in.shape[0], 1, 1, 1), device=x_in.device) - bias_final = [0.0] * x_in.shape[0] - - # perform calc_cond_uncond_batch per context window + # perform calc_conds_batch per context window for ctx_idxs in context_windows: ADGS.params.sub_idxs = ctx_idxs if ADGS.motion_models is not None: @@ -566,13 +565,12 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, for n in range(batched_conds): for ind in ctx_idxs: full_idxs.append((ADGS.params.full_length*n)+ind) - # get subsections of x, timestep, cond, uncond, cond_concat + # get subsections of x, timestep, conds sub_x = x_in[full_idxs] sub_timestep = timestep[full_idxs] - sub_cond = get_resized_cond(cond, full_idxs, len(ctx_idxs)) if cond is not None else None - sub_uncond = get_resized_cond(uncond, full_idxs, len(ctx_idxs)) if uncond is not None else None + sub_conds = [get_resized_cond(cond, full_idxs, len(ctx_idxs)) for cond in conds] - sub_cond_out, sub_uncond_out = calc_cond_uncond_batch_wrapper(model, [sub_cond, sub_uncond], sub_x, sub_timestep, model_options) + sub_conds_out = calc_cond_uncond_batch_wrapper(model, sub_conds, sub_x, sub_timestep, model_options) if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE: full_length = ADGS.params.full_length @@ -582,31 +580,31 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, bias = max(1e-2, bias) # take weighted average relative to total bias of current idx # and account for batched_conds - for n in range(batched_conds): - bias_total = bias_final[(full_length*n)+idx] - prev_weight = (bias_total / (bias_total + bias)) - new_weight = (bias / (bias_total + bias)) - cond_final[(full_length*n)+idx] = cond_final[(full_length*n)+idx] * prev_weight + sub_cond_out[(full_length*n)+pos] * new_weight - uncond_final[(full_length*n)+idx] = uncond_final[(full_length*n)+idx] * prev_weight + sub_uncond_out[(full_length*n)+pos] * new_weight - bias_final[(full_length*n)+idx] = bias_total + bias + for i in range(len(sub_conds_out)): + for n in range(batched_conds): + bias_total = biases_final[i][(full_length*n)+idx] + prev_weight = (bias_total / (bias_total + bias)) + new_weight = (bias / (bias_total + bias)) + conds_final[i][(full_length*n)+idx] = conds_final[i][(full_length*n)+idx] * prev_weight + sub_conds_out[i][(full_length*n)+pos] * new_weight + biases_final[i][(full_length*n)+idx] = bias_total + bias else: # add conds and counts based on weights of fuse method weights = get_context_weights(len(ctx_idxs), ADGS.params.context_options.fuse_method) * batched_conds weights_tensor = torch.Tensor(weights).to(device=x_in.device).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1) - cond_final[full_idxs] += sub_cond_out * weights_tensor - uncond_final[full_idxs] += sub_uncond_out * weights_tensor - out_count_final[full_idxs] += weights_tensor - - if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE: - # already normalized, so return as is - del out_count_final - return cond_final, uncond_final - else: - # normalize cond and uncond via division by context usage counts - cond_final /= out_count_final - uncond_final /= out_count_final - del out_count_final - return cond_final, uncond_final + for i in range(len(sub_conds_out)): + conds_final[i][full_idxs] += sub_conds_out[i] * weights_tensor + counts_final[i][full_idxs] += weights_tensor + + if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE: + # already normalized, so return as is + del counts_final + return conds_final + else: + # normalize conds via division by context usage counts + for i in range(len(conds_final)): + conds_final[i] /= counts_final[i] + del counts_final + return conds_final def calc_cond_uncond_batch_wrapper(model, conds: list[dict], x_in: Tensor, timestep, model_options): @@ -622,39 +620,34 @@ def calc_cond_uncond_batch_wrapper(model, conds: list[dict], x_in: Tensor, times if contains_lora_hooks: break if contains_lora_hooks: - return calc_cond_uncond_batch_lora_hook(model, conds[0], conds[1], x_in, timestep, model_options) + return calc_conds_batch_lora_hook(model, conds, x_in, timestep, model_options) # keep for backwards compatibility, for now if not hasattr(comfy.samplers, "calc_cond_batch"): return comfy.samplers.calc_cond_uncond_batch(model, conds[0], conds[1], x_in, timestep, model_options) return comfy.samplers.calc_cond_batch(model, conds, x_in, timestep, model_options) -def calc_cond_uncond_batch_lora_hook(model, cond, uncond, x_in, timestep, model_options): - out_cond = torch.zeros_like(x_in) - out_count = torch.ones_like(x_in) * 1e-37 - out_uncond = torch.zeros_like(x_in) - out_uncond_count = torch.ones_like(x_in) * 1e-37 +# based on comfy.samplers.calc_conds_batch +def calc_conds_batch_lora_hook(model: BaseModel, conds: list[dict], x_in: Tensor, timestep, model_options: dict): + out_conds = [] + out_counts = [] + # separate conds by matching lora_hooks + hooked_to_run: dict[LoraHookGroup,tuple[dict,int]] = {} - COND = 0 - UNCOND = 1 + # cond is i=0, uncond is i=1 + for i in range(len(conds)): + out_conds.append(torch.zeros_like(x_in)) + out_counts.append(torch.ones_like(x_in) * 1e-37) - # separate conds and unconds by matching lora_hooks - hooked_to_run = {} - for x in cond: - p = comfy.samplers.get_area_and_mult(x, x_in, timestep) - if p is None: - continue - hook: LoraHookGroup = x.get("lora_hook", None) - hooked_to_run.setdefault(hook, list()) - hooked_to_run[hook] += [(p, COND)] - if uncond is not None: - for x in uncond: - p = comfy.samplers.get_area_and_mult(x, x_in, timestep) - if p is None: - continue - hook: LoraHookGroup = x.get("lora_hook", None) - hooked_to_run.setdefault(hook, list()) - hooked_to_run[hook] += [(p, UNCOND)] + cond = conds[i] + if cond is not None: + for x in cond: + p = comfy.samplers.get_area_and_mult(x, x_in, timestep) + if p is None: + continue + hook: LoraHookGroup = x.get("lora_hook", None) + hooked_to_run.setdefault(hook, list()) + hooked_to_run[hook] += [(p, i)] # run every hooked_to_run separately for lora_hooks, to_run in hooked_to_run.items(): @@ -729,19 +722,13 @@ def calc_cond_uncond_batch_lora_hook(model, cond, uncond, x_in, timestep, model_ output = model_options['model_function_wrapper'](model.apply_model, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond}).chunk(batch_chunks) else: output = model.apply_model(input_x, timestep_, **c).chunk(batch_chunks) - del input_x for o in range(batch_chunks): - if cond_or_uncond[o] == COND: - out_cond[:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += output[o] * mult[o] - out_count[:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += mult[o] - else: - out_uncond[:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += output[o] * mult[o] - out_uncond_count[:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += mult[o] - del mult + cond_index = cond_or_uncond[o] + out_conds[cond_index][:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += output[o] * mult[o] + out_counts[cond_index][:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += mult[o] - out_cond /= out_count - del out_count - out_uncond /= out_uncond_count - del out_uncond_count - return out_cond, out_uncond + for i in range(len(out_conds)): + out_conds[i] /= out_counts[i] + + return out_conds From 44a097d803a4631f113fca2abe205455f7994966 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Thu, 25 Apr 2024 19:08:11 -0500 Subject: [PATCH 14/30] Added masked conditioning nodes to make masking/combining/applying hooks much easier, including Set Unmasked Conds with supporting code that makes specific conds apply to the areas where no masks were applied automatically --- animatediff/conditioning.py | 84 ++++++++++++++ animatediff/nodes.py | 15 ++- animatediff/nodes_conditioning.py | 177 ++++++++++++++++++++++++++++-- animatediff/sampling.py | 166 ++++++++++++++++++++++++++-- 4 files changed, 422 insertions(+), 20 deletions(-) diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py index e69de29..58d31c0 100644 --- a/animatediff/conditioning.py +++ b/animatediff/conditioning.py @@ -0,0 +1,84 @@ +from torch import Tensor + +from .utils_motion import LoraHookGroup + + +class COND_CONST: + KEY_LORA_HOOK = "lora_hook" + KEY_DEFAULT_COND = "default_cond" + + COND_AREA_DEFAULT = "default" + COND_AREA_MASK_BOUNDS = "mask bounds" + _LIST_COND_AREA = [COND_AREA_DEFAULT, COND_AREA_MASK_BOUNDS] + + +class ScheduleCond: + def __init__(self, start_percent: float, end_percent: float): + self.start_percent = start_percent + self.end_percent = end_percent + + +def conditioning_set_values(conditioning, values={}): + c = [] + for t in conditioning: + n = [t[0], t[1].copy()] + for k in values: + n[1][k] = values[k] + c.append(n) + return c + +def set_lora_hook_for_conditioning(conditioning, lora_hook: LoraHookGroup): + if lora_hook is None: + return conditioning + return conditioning_set_values(conditioning, {COND_CONST.KEY_LORA_HOOK: lora_hook}) + +def set_schedule_for_conditioning(conditioning, start_percent: float, end_percent: float): + return conditioning_set_values(conditioning, {"start_percent": start_percent, "end_percent": end_percent}) + +def set_mask_for_conditioning(conditioning, mask: Tensor, set_cond_area: str, strength: float): + set_area_to_bounds = False + if set_cond_area != COND_CONST.COND_AREA_DEFAULT: + set_area_to_bounds = True + if len(mask.shape) < 3: + mask = mask.unsqueeze(0) + + return conditioning_set_values(conditioning, {"mask": mask, + "set_area_to_bounds": set_area_to_bounds, + "mask_strength": strength}) + +def combine_conditioning(conds: list): + combined_conds = [] + for cond in conds: + combined_conds.extend(cond) + return combined_conds + +def set_mask_conds(conds: list, mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + masked_conds = [] + for c in conds: + # first, apply lora_hook to conditioning, if provided + c = set_lora_hook_for_conditioning(c, opt_lora_hook) + # finally, apply mask to conditioning and store + masked_conds.append(set_mask_for_conditioning(conditioning=c, mask=mask, strength=strength, set_cond_area=set_cond_area)) + return masked_conds + +def set_mask_and_combine_conds(conds: list, new_conds: list, mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + combined_conds = [] + for c, masked_c in zip(conds, new_conds): + # first, apply lora_hook to new conditioning, if provided + masked_c = set_lora_hook_for_conditioning(masked_c, opt_lora_hook) + # next, apply mask to new conditioning + masked_c = set_mask_for_conditioning(conditioning=masked_c, mask=mask, set_cond_area=set_cond_area, strength=strength) + # finally, combine with existing conditioning and store + combined_conds.append(combine_conditioning([c, masked_c])) + return combined_conds + +def set_unmasked_and_combine_conds(conds: list, new_conds: list, opt_lora_hook: LoraHookGroup): + combined_conds = [] + for c, new_c in zip(conds, new_conds): + # first, apply lora_hook to new conditioning, if provided + new_c = set_lora_hook_for_conditioning(new_c, opt_lora_hook) + # next, add default_cond key to cond so that during sampling, it can be identified + new_c = conditioning_set_values(new_c, {COND_CONST.KEY_DEFAULT_COND: True}) + # finally, combine with existing conditioning and store + combined_conds.append(combine_conditioning([c, new_c])) + return combined_conds diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 306cca4..83b4df6 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -12,7 +12,10 @@ from .nodes_cameractrl import (LoadAnimateDiffModelWithCameraCtrl, ApplyAnimateD from .nodes_multival import MultivalDynamicNode, MultivalScaledMaskNode from .nodes_conditioning import (MaskableLoraLoader, MaskableLoraLoaderModelOnly, MaskableSDModelLoader, MaskableSDModelLoaderModelOnly, SetModelLoraHook, SetClipLoraHook, - CombineLoraHooks, CombineLoraHookFourOptional, CombineLoraHookEightOptional) + CombineLoraHooks, CombineLoraHookFourOptional, CombineLoraHookEightOptional, + PairedConditioningSetMaskHooked, ConditioningSetMaskHooked, + PairedConditioningSetMaskAndCombineHooked, ConditioningSetMaskAndCombineHooked, + PairedConditioningSetUnmaskedAndCombineHooked) from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGKeyframeNode) from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode) @@ -65,6 +68,11 @@ NODE_CLASS_MAPPINGS = { "ADE_CombineLoraHooksEight": CombineLoraHookEightOptional, "ADE_AttachLoraHookToConditioning": SetModelLoraHook, "ADE_AttachLoraHookToCLIP": SetClipLoraHook, + "ADE_PairedConditioningSetMask": PairedConditioningSetMaskHooked, + "ADE_ConditioningSetMask": ConditioningSetMaskHooked, + "ADE_PairedConditioningSetMaskAndCombine": PairedConditioningSetMaskAndCombineHooked, + "ADE_ConditioningSetMaskAndCombine": ConditioningSetMaskAndCombineHooked, + "ADE_PairedConditioningSetUnmaskedAndCombine": PairedConditioningSetUnmaskedAndCombineHooked, # Noise Layer Nodes "ADE_NoiseLayerAdd": NoiseLayerAddNode, "ADE_NoiseLayerAddWeighted": NoiseLayerAddWeightedNode, @@ -159,6 +167,11 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_CombineLoraHooksEight": "Combine LoRA Hooks [8] πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", + "ADE_PairedConditioningSetMask": "Set Mask on Conds πŸŽ­πŸ…πŸ…“", + "ADE_ConditioningSetMask": "Set Mask on Cond πŸŽ­πŸ…πŸ…“", + "ADE_PairedConditioningSetMaskAndCombine": "Set Mask and Combine Conds πŸŽ­πŸ…πŸ…“", + "ADE_ConditioningSetMaskAndCombine": "Set Mask and Combine Cond πŸŽ­πŸ…πŸ…“", + "ADE_PairedConditioningSetUnmaskedAndCombine": "Set Unmasked Conds πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index d759fb7..1b28de8 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -1,15 +1,166 @@ import uuid import folder_paths from typing import Union +from torch import Tensor from comfy.model_patcher import ModelPatcher from comfy.sd import CLIP import comfy.sd import comfy.utils -from .utils_motion import LoraHook, LoraHookGroup +from .conditioning import (COND_CONST, set_mask_conds, set_mask_and_combine_conds, set_unmasked_and_combine_conds, + set_lora_hook_for_conditioning, combine_conditioning) from .model_injection import ModelPatcherAndInjector, CLIPWithHooks, load_hooked_lora_for_models, load_model_as_hooked_lora_for_models +from .utils_motion import LoraHook, LoraHookGroup + +############################################### +### Mask, Combine, and Hook Conditioning +############################################### +class PairedConditioningSetMaskHooked: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "positive_ADD": ("CONDITIONING", ), + "negative_ADD": ("CONDITIONING", ), + "mask": ("MASK", ), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "set_cond_area": (COND_CONST._LIST_COND_AREA,), + }, + "optional": { + "opt_lora_hook": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("CONDITIONING", "CONDITIONING") + RETURN_NAMES = ("positive", "negative") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "append_and_hook" + + def append_and_hook(self, positive_ADD, negative_ADD, + mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + final_positive, final_negative = set_mask_conds(conds=[positive_ADD, negative_ADD], + mask=mask, strength=strength, set_cond_area=set_cond_area, opt_lora_hook=opt_lora_hook) + return (final_positive, final_negative) + + +class ConditioningSetMaskHooked: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cond_ADD": ("CONDITIONING",), + "mask": ("MASK",), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "set_cond_area": (COND_CONST._LIST_COND_AREA,), + }, + "optional": { + "opt_lora_hook": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("CONDITIONING",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "append_and_hook" + + def append_and_hook(self, cond_ADD, + mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + (final_conditioning,) = set_mask_conds(conds=[cond_ADD], + mask=mask, strength=strength, set_cond_area=set_cond_area, opt_lora_hook=opt_lora_hook) + return (final_conditioning,) + + +class PairedConditioningSetMaskAndCombineHooked: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "positive_ADD": ("CONDITIONING",), + "negative_ADD": ("CONDITIONING",), + "mask": ("MASK",), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "set_cond_area": (COND_CONST._LIST_COND_AREA,), + }, + "optional": { + "opt_lora_hook": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("CONDITIONING", "CONDITIONING") + RETURN_NAMES = ("positive", "negative") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "append_and_combine" + + def append_and_combine(self, positive, negative, positive_ADD, negative_ADD, + mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + final_positive, final_negative = set_mask_and_combine_conds(conds=[positive, negative], new_conds=[positive_ADD, negative_ADD], + mask=mask, strength=strength, set_cond_area=set_cond_area, opt_lora_hook=opt_lora_hook) + return (final_positive, final_negative,) + + +class ConditioningSetMaskAndCombineHooked: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cond": ("CONDITIONING",), + "cond_ADD": ("CONDITIONING",), + "mask": ("MASK", ), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "set_cond_area": (COND_CONST._LIST_COND_AREA,), + }, + "optional": { + "opt_lora_hook": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("CONDITIONING",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "append_and_combine" + + def append_and_combine(self, conditioning, conditioning_ADD, + mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + (final_conditioning,) = set_mask_and_combine_conds(conds=[conditioning], new_conds=[conditioning_ADD], + mask=mask, strength=strength, set_cond_area=set_cond_area, opt_lora_hook=opt_lora_hook) + return (final_conditioning,) + + +class PairedConditioningSetUnmaskedAndCombineHooked: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "positive_DEFAULT": ("CONDITIONING",), + "negative_DEFAULT": ("CONDITIONING",), + }, + "optional": { + "opt_lora_hook": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("CONDITIONING", "CONDITIONING") + RETURN_NAMES = ("positive", "negative") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "append_and_combine" + + def append_and_combine(self, positive, negative, positive_DEFAULT, negative_DEFAULT, + opt_lora_hook: LoraHookGroup=None): + final_positive, final_negative = set_unmasked_and_combine_conds(conds=[positive, negative], new_conds=[positive_DEFAULT, negative_DEFAULT], + opt_lora_hook=opt_lora_hook) + return (final_positive, final_negative,) +############################################### +############################################### +############################################### + + +############################################### +### Register LoRA Hooks +############################################### # based on ComfyUI's nodes.py LoraLoader class MaskableLoraLoader: def __init__(self): @@ -28,7 +179,7 @@ class MaskableLoraLoader: } RETURN_TYPES = ("MODEL", "CLIP", "LORA_HOOK") - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/register lora hooks" FUNCTION = "load_lora" def load_lora(self, model: Union[ModelPatcher, ModelPatcherAndInjector], clip: CLIP, lora_name: str, strength_model: float, strength_clip: float): @@ -69,7 +220,7 @@ class MaskableLoraLoaderModelOnly(MaskableLoraLoader): } RETURN_TYPES = ("MODEL", "LORA_HOOK") - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/register lora hooks" FUNCTION = "load_lora_model_only" def load_lora_model_only(self, model: ModelPatcher, lora_name: str, strength_model: float): @@ -90,7 +241,7 @@ class MaskableSDModelLoader: } RETURN_TYPES = ("MODEL", "CLIP", "LORA_HOOK") - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/register lora hooks" FUNCTION = "load_model_as_lora" def load_model_as_lora(self, model: ModelPatcher, clip: CLIP, ckpt_name: str): @@ -119,14 +270,21 @@ class MaskableSDModelLoaderModelOnly(MaskableSDModelLoader): } RETURN_TYPES = ("MODEL", "LORA_HOOK") - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/register lora hooks" FUNCTION = "load_model_as_lora_model_only" def load_model_as_lora_model_only(self, model: ModelPatcher, ckpt_name: str): model_lora, clip_lora, lora_hook = self.load_model_as_lora(model=model, clip=None, ckpt_name=ckpt_name) return (model_lora, lora_hook) +############################################### +############################################### +############################################### + +############################################### +### Set LoRA Hooks +############################################### class SetModelLoraHook: @classmethod def INPUT_TYPES(s): @@ -182,7 +340,7 @@ class CombineLoraHooks: } RETURN_TYPES = ("LORA_HOOK",) - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine hooks" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine lora hooks" FUNCTION = "combine_lora_hooks" def combine_lora_hooks(self, lora_hook_A: LoraHookGroup, lora_hook_B: LoraHookGroup): @@ -204,7 +362,7 @@ class CombineLoraHookFourOptional: } RETURN_TYPES = ("LORA_HOOK",) - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine hooks" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine lora hooks" FUNCTION = "combine_lora_hooks" def combine_lora_hooks(self, @@ -233,7 +391,7 @@ class CombineLoraHookEightOptional: } RETURN_TYPES = ("LORA_HOOK",) - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine hooks" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine lora hooks" FUNCTION = "combine_lora_hooks" def combine_lora_hooks(self, @@ -247,3 +405,6 @@ class CombineLoraHookEightOptional: # NOTE: if at some point I add more Javascript stuff to this repo, there should be a combine node # that dynamically increases the hooks available to plug in on the node +############################################### +############################################### +############################################### diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 8160939..595cc44 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -1,5 +1,6 @@ from typing import Callable +import collections import math import torch from torch import Tensor @@ -21,6 +22,7 @@ from comfy.controlnet import ControlBase from comfy.model_base import BaseModel import comfy.ops +from .conditioning import COND_CONST from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows from .sample_settings import IterationOptions, SampleSettings, SeedNoiseGeneration from .utils_model import ModelTypeSD @@ -221,12 +223,13 @@ class FunctionInjectionHolder: pass def inject_functions(self, model: ModelPatcherAndInjector, params: InjectionParams): - # Save Original Functions + # Save Original Functions - order must match between here and restore_functions self.orig_forward_timestep_embed = openaimodel.forward_timestep_embed # needed to account for VanillaTemporalModule self.orig_memory_required = model.model.memory_required # allows for "unlimited area hack" to prevent halving of conds/unconds self.orig_groupnorm_forward = torch.nn.GroupNorm.forward # used to normalize latents to remove "flickering" of colors/brightness between frames self.orig_groupnorm_manual_cast_forward = comfy.ops.manual_cast.GroupNorm.forward_comfy_cast_weights self.orig_sampling_function = comfy.samplers.sampling_function # used to support sliding context windows in samplers + self.orig_get_area_and_mult = comfy.samplers.get_area_and_mult if SAMPLE_FALLBACK: # for backwards compatibility, for now self.orig_get_additional_models = comfy.sample.get_additional_models else: @@ -256,6 +259,7 @@ class FunctionInjectionHolder: break del info comfy.samplers.sampling_function = evolved_sampling_function + comfy.samplers.get_area_and_mult = get_area_and_mult_ADE if SAMPLE_FALLBACK: # for backwards compatibility, for now comfy.sample.get_additional_models = get_additional_models_factory(self.orig_get_additional_models, model.motion_models) else: @@ -268,6 +272,7 @@ class FunctionInjectionHolder: openaimodel.forward_timestep_embed = self.orig_forward_timestep_embed torch.nn.GroupNorm.forward = self.orig_groupnorm_forward comfy.ops.manual_cast.GroupNorm.forward_comfy_cast_weights = self.orig_groupnorm_manual_cast_forward + comfy.samplers.get_area_and_mult = self.orig_get_area_and_mult comfy.samplers.sampling_function = self.orig_sampling_function if SAMPLE_FALLBACK: # for backwards compatibility, for now comfy.sample.get_additional_models = self.orig_get_additional_models @@ -608,31 +613,166 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options def calc_cond_uncond_batch_wrapper(model, conds: list[dict], x_in: Tensor, timestep, model_options): - # check if conds or unconds contain lora_hook + # check if conds or unconds contain lora_hook or default_cond contains_lora_hooks = False + has_default_cond = False for cond_uncond in conds: if cond_uncond is None: continue for t in cond_uncond: - if "lora_hook" in t: + if COND_CONST.KEY_LORA_HOOK in t: contains_lora_hooks = True - break - if contains_lora_hooks: - break - if contains_lora_hooks: - return calc_conds_batch_lora_hook(model, conds, x_in, timestep, model_options) + if COND_CONST.KEY_DEFAULT_COND in t: + has_default_cond = True + # if contains_lora_hooks: + # break + if contains_lora_hooks or has_default_cond: + return calc_conds_batch_lora_hook(model, conds, x_in, timestep, model_options, has_default_cond) # keep for backwards compatibility, for now if not hasattr(comfy.samplers, "calc_cond_batch"): return comfy.samplers.calc_cond_uncond_batch(model, conds[0], conds[1], x_in, timestep, model_options) return comfy.samplers.calc_cond_batch(model, conds, x_in, timestep, model_options) +# modified from comfy.samplers.get_area_and_mult +COND_OBJ = collections.namedtuple('cond_obj', ['input_x', 'mult', 'conditioning', 'area', 'control', 'patches']) +def get_area_and_mult_ADE(conds, x_in, timestep_in): + area = (x_in.shape[2], x_in.shape[3], 0, 0) + strength = 1.0 + + if 'timestep_start' in conds: + timestep_start = conds['timestep_start'] + if timestep_in[0] > timestep_start: + return None + if 'timestep_end' in conds: + timestep_end = conds['timestep_end'] + if timestep_in[0] < timestep_end: + return None + if 'area' in conds: + area = conds['area'] + if 'strength' in conds: + strength = conds['strength'] + + input_x = x_in[:,:,area[2]:area[0] + area[2],area[3]:area[1] + area[3]] + if 'mask' in conds: + # Scale the mask to the size of the input + # The mask should have been resized as we began the sampling process + mask_strength = 1.0 + if "mask_strength" in conds: + mask_strength = conds["mask_strength"] + mask = conds['mask'] + assert(mask.shape[1] == x_in.shape[2]) + assert(mask.shape[2] == x_in.shape[3]) + # make sure mask is capped at input_shape batch length to prevent 0 as dimension + mask = mask[:input_x.shape[0], area[2]:area[0] + area[2], area[3]:area[1] + area[3]] * mask_strength + mask = mask.unsqueeze(1).repeat(input_x.shape[0] // mask.shape[0], input_x.shape[1], 1, 1) + else: + mask = torch.ones_like(input_x) + mult = mask * strength + + if 'mask' not in conds: + rr = 8 + if area[2] != 0: + for t in range(rr): + mult[:,:,t:1+t,:] *= ((1.0/rr) * (t + 1)) + if (area[0] + area[2]) < x_in.shape[2]: + for t in range(rr): + mult[:,:,area[0] - 1 - t:area[0] - t,:] *= ((1.0/rr) * (t + 1)) + if area[3] != 0: + for t in range(rr): + mult[:,:,:,t:1+t] *= ((1.0/rr) * (t + 1)) + if (area[1] + area[3]) < x_in.shape[3]: + for t in range(rr): + mult[:,:,:,area[1] - 1 - t:area[1] - t] *= ((1.0/rr) * (t + 1)) + + conditioning = {} + model_conds = conds["model_conds"] + for c in model_conds: + conditioning[c] = model_conds[c].process_cond(batch_size=x_in.shape[0], device=x_in.device, area=area) + + control = conds.get('control', None) + + patches = None + if 'gligen' in conds: + gligen = conds['gligen'] + patches = {} + gligen_type = gligen[0] + gligen_model = gligen[1] + if gligen_type == "position": + gligen_patch = gligen_model.model.set_position(input_x.shape, gligen[2], input_x.device) + else: + gligen_patch = gligen_model.model.set_empty(input_x.shape, input_x.device) + + patches['middle_patch'] = [gligen_patch] + + return COND_OBJ(input_x, mult, conditioning, area, control, patches) + + +def separate_default_conds(conds: list[dict]): + normal_conds = [] + default_conds = [] + for i in range(len(conds)): + c = [] + default_c = [] + for t in conds[i]: + # check if cond is a default cond + if COND_CONST.KEY_DEFAULT_COND in t: + default_c.append(t) + else: + c.append(t) + normal_conds.append(c) + default_conds.append(default_c) + return normal_conds, default_conds + + +def finalize_default_conds(hooked_to_run: dict[LoraHookGroup,list[tuple[COND_OBJ,int]]], default_conds: list[list[dict]], x_in: Tensor, timestep): + # need to figure out remaining unmasked area for conds + default_mults = [] + for d in default_conds: + default_mults.append(torch.ones_like(x_in)) + + # look through each finalized cond in hooked_to_run for 'mult' and subtract it from each cond + for lora_hooks, to_run in hooked_to_run.items(): + for cond_obj, i in to_run: + # if no default_cond for cond_type, do nothing + if len(default_conds[i]) == 0: + continue + area: list[int] = cond_obj.area + default_mults[i][:,:,area[2]:area[0] + area[2],area[3]:area[1] + area[3]] -= cond_obj.mult + + # for each default_mult, ReLU to make negatives=0, and then check for any nonzeros + for i, mult in enumerate(default_mults): + # if no default_cond for cond type, do nothing + if len(default_conds[i]) == 0: + continue + torch.nn.functional.relu(mult, inplace=True) + # if mult is all zeros, then don't add default_cond + if torch.max(mult) == 0.0: + continue + + cond = default_conds[i] + for x in cond: + # do get_area_and_mult to get all the expected values + p = comfy.samplers.get_area_and_mult(x, x_in, timestep) + if p is None: + continue + # replace p's mult with calculated mult + p = p._replace(mult=mult) + hook: LoraHookGroup = x.get(COND_CONST.KEY_LORA_HOOK, None) + hooked_to_run.setdefault(hook, list()) + hooked_to_run[hook] += [(p, i)] + + # based on comfy.samplers.calc_conds_batch -def calc_conds_batch_lora_hook(model: BaseModel, conds: list[dict], x_in: Tensor, timestep, model_options: dict): +def calc_conds_batch_lora_hook(model: BaseModel, conds: list[list[dict]], x_in: Tensor, timestep, model_options: dict, has_default_cond=False): out_conds = [] out_counts = [] # separate conds by matching lora_hooks - hooked_to_run: dict[LoraHookGroup,tuple[dict,int]] = {} + hooked_to_run: dict[LoraHookGroup,list[tuple[collections.namedtuple,int]]] = {} + + # separate out default_conds, if needed + if has_default_cond: + conds, default_conds = separate_default_conds(conds) # cond is i=0, uncond is i=1 for i in range(len(conds)): @@ -645,10 +785,14 @@ def calc_conds_batch_lora_hook(model: BaseModel, conds: list[dict], x_in: Tensor p = comfy.samplers.get_area_and_mult(x, x_in, timestep) if p is None: continue - hook: LoraHookGroup = x.get("lora_hook", None) + hook: LoraHookGroup = x.get(COND_CONST.KEY_LORA_HOOK, None) hooked_to_run.setdefault(hook, list()) hooked_to_run[hook] += [(p, i)] + # finalize default_conds, if needed + if has_default_cond: + finalize_default_conds(hooked_to_run, default_conds, x_in, timestep) + # run every hooked_to_run separately for lora_hooks, to_run in hooked_to_run.items(): while len(to_run) > 0: From eaf58bc46006b8722a80a21b2eb7c952e5de5b37 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 26 Apr 2024 07:41:43 -0500 Subject: [PATCH 15/30] Fixed sliding context - accidentally indented the combining code when refactored sliding_calc_conds_batch, making the function return immediately after the first context is sampled --- animatediff/sampling.py | 20 ++++++++++---------- 1 file changed, 10 insertions(+), 10 deletions(-) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 595cc44..05b3504 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -600,16 +600,16 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options conds_final[i][full_idxs] += sub_conds_out[i] * weights_tensor counts_final[i][full_idxs] += weights_tensor - if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE: - # already normalized, so return as is - del counts_final - return conds_final - else: - # normalize conds via division by context usage counts - for i in range(len(conds_final)): - conds_final[i] /= counts_final[i] - del counts_final - return conds_final + if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE: + # already normalized, so return as is + del counts_final + return conds_final + else: + # normalize conds via division by context usage counts + for i in range(len(conds_final)): + conds_final[i] /= counts_final[i] + del counts_final + return conds_final def calc_cond_uncond_batch_wrapper(model, conds: list[dict], x_in: Tensor, timestep, model_options): From 2b6c1bf4753a7c037b9fa752360440d20dabd516 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 26 Apr 2024 10:27:10 -0500 Subject: [PATCH 16/30] Added support for Timesteps Conditioning, masks are now optional (so the other parts of the nodes can be used without masks), renaming and reorganizing nodes --- animatediff/conditioning.py | 34 ++++++++--- animatediff/nodes.py | 15 +++-- animatediff/nodes_conditioning.py | 94 +++++++++++++++++++++++++------ 3 files changed, 112 insertions(+), 31 deletions(-) diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py index 58d31c0..7c0c39b 100644 --- a/animatediff/conditioning.py +++ b/animatediff/conditioning.py @@ -12,7 +12,7 @@ class COND_CONST: _LIST_COND_AREA = [COND_AREA_DEFAULT, COND_AREA_MASK_BOUNDS] -class ScheduleCond: +class TimestepsCond: def __init__(self, start_percent: float, end_percent: float): self.start_percent = start_percent self.end_percent = end_percent @@ -32,10 +32,15 @@ def set_lora_hook_for_conditioning(conditioning, lora_hook: LoraHookGroup): return conditioning return conditioning_set_values(conditioning, {COND_CONST.KEY_LORA_HOOK: lora_hook}) -def set_schedule_for_conditioning(conditioning, start_percent: float, end_percent: float): - return conditioning_set_values(conditioning, {"start_percent": start_percent, "end_percent": end_percent}) +def set_timesteps_for_conditioning(conditioning, timesteps_cond: TimestepsCond): + if timesteps_cond is None: + return conditioning + return conditioning_set_values(conditioning, {"start_percent": timesteps_cond.start_percent, + "end_percent": timesteps_cond.end_percent}) def set_mask_for_conditioning(conditioning, mask: Tensor, set_cond_area: str, strength: float): + if mask is None: + return conditioning set_area_to_bounds = False if set_cond_area != COND_CONST.COND_AREA_DEFAULT: set_area_to_bounds = True @@ -52,33 +57,44 @@ def combine_conditioning(conds: list): combined_conds.extend(cond) return combined_conds -def set_mask_conds(conds: list, mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): +def set_mask_conds(conds: list, strength: float, set_cond_area: str, + opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None): masked_conds = [] for c in conds: # first, apply lora_hook to conditioning, if provided c = set_lora_hook_for_conditioning(c, opt_lora_hook) + # next, apply mask to conditioning + c = set_mask_for_conditioning(conditioning=c, mask=opt_mask, strength=strength, set_cond_area=set_cond_area) + # apply timesteps, if present + c = set_timesteps_for_conditioning(conditioning=c, timesteps_cond=opt_timesteps) # finally, apply mask to conditioning and store - masked_conds.append(set_mask_for_conditioning(conditioning=c, mask=mask, strength=strength, set_cond_area=set_cond_area)) + masked_conds.append(c) return masked_conds -def set_mask_and_combine_conds(conds: list, new_conds: list, mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): +def set_mask_and_combine_conds(conds: list, new_conds: list, strength: float, set_cond_area: str, + opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None): combined_conds = [] for c, masked_c in zip(conds, new_conds): # first, apply lora_hook to new conditioning, if provided masked_c = set_lora_hook_for_conditioning(masked_c, opt_lora_hook) - # next, apply mask to new conditioning - masked_c = set_mask_for_conditioning(conditioning=masked_c, mask=mask, set_cond_area=set_cond_area, strength=strength) + # next, apply mask to new conditioning, if provided + masked_c = set_mask_for_conditioning(conditioning=masked_c, mask=opt_mask, set_cond_area=set_cond_area, strength=strength) + # apply timesteps, if present + masked_c = set_timesteps_for_conditioning(conditioning=masked_c, timesteps_cond=opt_timesteps) # finally, combine with existing conditioning and store combined_conds.append(combine_conditioning([c, masked_c])) return combined_conds -def set_unmasked_and_combine_conds(conds: list, new_conds: list, opt_lora_hook: LoraHookGroup): +def set_unmasked_and_combine_conds(conds: list, new_conds: list, + opt_lora_hook: LoraHookGroup, opt_timesteps: TimestepsCond=None): combined_conds = [] for c, new_c in zip(conds, new_conds): # first, apply lora_hook to new conditioning, if provided new_c = set_lora_hook_for_conditioning(new_c, opt_lora_hook) # next, add default_cond key to cond so that during sampling, it can be identified new_c = conditioning_set_values(new_c, {COND_CONST.KEY_DEFAULT_COND: True}) + # apply timesteps, if present + new_c = set_timesteps_for_conditioning(conditioning=new_c, timesteps_cond=opt_timesteps) # finally, combine with existing conditioning and store combined_conds.append(combine_conditioning([c, new_c])) return combined_conds diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 83b4df6..027345c 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -15,7 +15,8 @@ from .nodes_conditioning import (MaskableLoraLoader, MaskableLoraLoaderModelOnly CombineLoraHooks, CombineLoraHookFourOptional, CombineLoraHookEightOptional, PairedConditioningSetMaskHooked, ConditioningSetMaskHooked, PairedConditioningSetMaskAndCombineHooked, ConditioningSetMaskAndCombineHooked, - PairedConditioningSetUnmaskedAndCombineHooked) + PairedConditioningSetUnmaskedAndCombineHooked, ConditioningSetUnmaskedAndCombineHooked, + ConditioningTimestepsNode) from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGKeyframeNode) from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode) @@ -73,6 +74,8 @@ NODE_CLASS_MAPPINGS = { "ADE_PairedConditioningSetMaskAndCombine": PairedConditioningSetMaskAndCombineHooked, "ADE_ConditioningSetMaskAndCombine": ConditioningSetMaskAndCombineHooked, "ADE_PairedConditioningSetUnmaskedAndCombine": PairedConditioningSetUnmaskedAndCombineHooked, + "ADE_ConditioningSetUnmaskedAndCombine": ConditioningSetUnmaskedAndCombineHooked, + "ADE_TimestepsConditioning": ConditioningTimestepsNode, # Noise Layer Nodes "ADE_NoiseLayerAdd": NoiseLayerAddNode, "ADE_NoiseLayerAddWeighted": NoiseLayerAddWeightedNode, @@ -167,11 +170,13 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_CombineLoraHooksEight": "Combine LoRA Hooks [8] πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", - "ADE_PairedConditioningSetMask": "Set Mask on Conds πŸŽ­πŸ…πŸ…“", - "ADE_ConditioningSetMask": "Set Mask on Cond πŸŽ­πŸ…πŸ…“", - "ADE_PairedConditioningSetMaskAndCombine": "Set Mask and Combine Conds πŸŽ­πŸ…πŸ…“", - "ADE_ConditioningSetMaskAndCombine": "Set Mask and Combine Cond πŸŽ­πŸ…πŸ…“", + "ADE_PairedConditioningSetMask": "Set Props on Conds πŸŽ­πŸ…πŸ…“", + "ADE_ConditioningSetMask": "Set Props on Cond πŸŽ­πŸ…πŸ…“", + "ADE_PairedConditioningSetMaskAndCombine": "Set Props and Combine Conds πŸŽ­πŸ…πŸ…“", + "ADE_ConditioningSetMaskAndCombine": "Set Props and Combine Cond πŸŽ­πŸ…πŸ…“", "ADE_PairedConditioningSetUnmaskedAndCombine": "Set Unmasked Conds πŸŽ­πŸ…πŸ…“", + "ADE_ConditioningSetUnmaskedAndCombine": "Set Unmasked Cond πŸŽ­πŸ…πŸ…“", + "ADE_TimestepsConditioning": "Timesteps Conditioning πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 1b28de8..6150fec 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -8,8 +8,7 @@ from comfy.sd import CLIP import comfy.sd import comfy.utils -from .conditioning import (COND_CONST, set_mask_conds, set_mask_and_combine_conds, set_unmasked_and_combine_conds, - set_lora_hook_for_conditioning, combine_conditioning) +from .conditioning import (COND_CONST, TimestepsCond, set_mask_conds, set_mask_and_combine_conds, set_unmasked_and_combine_conds,) from .model_injection import ModelPatcherAndInjector, CLIPWithHooks, load_hooked_lora_for_models, load_model_as_hooked_lora_for_models from .utils_motion import LoraHook, LoraHookGroup @@ -24,12 +23,13 @@ class PairedConditioningSetMaskHooked: "required": { "positive_ADD": ("CONDITIONING", ), "negative_ADD": ("CONDITIONING", ), - "mask": ("MASK", ), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), "set_cond_area": (COND_CONST._LIST_COND_AREA,), }, "optional": { + "opt_mask": ("MASK", ), "opt_lora_hook": ("LORA_HOOK",), + "opt_timesteps": ("TIMESTEPS_COND",) } } @@ -39,9 +39,11 @@ class PairedConditioningSetMaskHooked: FUNCTION = "append_and_hook" def append_and_hook(self, positive_ADD, negative_ADD, - mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + strength: float, set_cond_area: str, + opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None): final_positive, final_negative = set_mask_conds(conds=[positive_ADD, negative_ADD], - mask=mask, strength=strength, set_cond_area=set_cond_area, opt_lora_hook=opt_lora_hook) + strength=strength, set_cond_area=set_cond_area, + opt_mask=opt_mask, opt_lora_hook=opt_lora_hook, opt_timesteps=opt_timesteps) return (final_positive, final_negative) @@ -51,23 +53,26 @@ class ConditioningSetMaskHooked: return { "required": { "cond_ADD": ("CONDITIONING",), - "mask": ("MASK",), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), "set_cond_area": (COND_CONST._LIST_COND_AREA,), }, "optional": { + "opt_mask": ("MASK", ), "opt_lora_hook": ("LORA_HOOK",), + "opt_timesteps": ("TIMESTEPS_COND",) } } RETURN_TYPES = ("CONDITIONING",) - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/single cond ops" FUNCTION = "append_and_hook" def append_and_hook(self, cond_ADD, - mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + strength: float, set_cond_area: str, + opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None): (final_conditioning,) = set_mask_conds(conds=[cond_ADD], - mask=mask, strength=strength, set_cond_area=set_cond_area, opt_lora_hook=opt_lora_hook) + strength=strength, set_cond_area=set_cond_area, + opt_mask=opt_mask, opt_lora_hook=opt_lora_hook, opt_timesteps=opt_timesteps) return (final_conditioning,) @@ -80,12 +85,13 @@ class PairedConditioningSetMaskAndCombineHooked: "negative": ("CONDITIONING",), "positive_ADD": ("CONDITIONING",), "negative_ADD": ("CONDITIONING",), - "mask": ("MASK",), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), "set_cond_area": (COND_CONST._LIST_COND_AREA,), }, "optional": { + "opt_mask": ("MASK", ), "opt_lora_hook": ("LORA_HOOK",), + "opt_timesteps": ("TIMESTEPS_COND",) } } @@ -95,9 +101,11 @@ class PairedConditioningSetMaskAndCombineHooked: FUNCTION = "append_and_combine" def append_and_combine(self, positive, negative, positive_ADD, negative_ADD, - mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + strength: float, set_cond_area: str, + opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None): final_positive, final_negative = set_mask_and_combine_conds(conds=[positive, negative], new_conds=[positive_ADD, negative_ADD], - mask=mask, strength=strength, set_cond_area=set_cond_area, opt_lora_hook=opt_lora_hook) + strength=strength, set_cond_area=set_cond_area, + opt_mask=opt_mask, opt_lora_hook=opt_lora_hook, opt_timesteps=opt_timesteps) return (final_positive, final_negative,) @@ -108,23 +116,26 @@ class ConditioningSetMaskAndCombineHooked: "required": { "cond": ("CONDITIONING",), "cond_ADD": ("CONDITIONING",), - "mask": ("MASK", ), "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), "set_cond_area": (COND_CONST._LIST_COND_AREA,), }, "optional": { + "opt_mask": ("MASK", ), "opt_lora_hook": ("LORA_HOOK",), + "opt_timesteps": ("TIMESTEPS_COND",) } } RETURN_TYPES = ("CONDITIONING",) - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/single cond ops" FUNCTION = "append_and_combine" def append_and_combine(self, conditioning, conditioning_ADD, - mask: Tensor, strength: float, set_cond_area: str, opt_lora_hook: LoraHookGroup=None): + strength: float, set_cond_area: str, + opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None): (final_conditioning,) = set_mask_and_combine_conds(conds=[conditioning], new_conds=[conditioning_ADD], - mask=mask, strength=strength, set_cond_area=set_cond_area, opt_lora_hook=opt_lora_hook) + strength=strength, set_cond_area=set_cond_area, + opt_mask=opt_mask, opt_lora_hook=opt_lora_hook, opt_timesteps=opt_timesteps) return (final_conditioning,) @@ -153,6 +164,55 @@ class PairedConditioningSetUnmaskedAndCombineHooked: final_positive, final_negative = set_unmasked_and_combine_conds(conds=[positive, negative], new_conds=[positive_DEFAULT, negative_DEFAULT], opt_lora_hook=opt_lora_hook) return (final_positive, final_negative,) + + +class ConditioningSetUnmaskedAndCombineHooked: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "cond": ("CONDITIONING",), + "cond_DEFAULT": ("CONDITIONING",), + }, + "optional": { + "opt_lora_hook": ("LORA_HOOK",), + } + } + + RETURN_TYPES = ("CONDITIONING",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/single cond ops" + FUNCTION = "append_and_combine" + + def append_and_combine(self, cond, cond_DEFAULT, + opt_lora_hook: LoraHookGroup=None): + (final_conditioning,) = set_unmasked_and_combine_conds(conds=[cond], new_conds=[cond_DEFAULT], + opt_lora_hook=opt_lora_hook) + return (final_conditioning,) +############################################### +############################################### +############################################### + + +############################################### +### Scheduling +############################################### +class ConditioningTimestepsNode: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}) + } + } + + RETURN_TYPES = ("TIMESTEPS_COND",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "create_schedule" + + def create_schedule(self, start_percent: float, end_percent: float): + return (TimestepsCond(start_percent=start_percent, end_percent=end_percent),) + ############################################### ############################################### ############################################### @@ -296,7 +356,7 @@ class SetModelLoraHook: } RETURN_TYPES = ("CONDITIONING",) - CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/single cond ops" FUNCTION = "attach_lora_hook" def attach_lora_hook(self, conditioning, lora_hook: LoraHookGroup): From 667042b43e9beef2cb75ef70d677e164b7b66d5d Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 26 Apr 2024 17:24:19 -0500 Subject: [PATCH 17/30] Initial lowvram support for lora hooks --- animatediff/model_injection.py | 91 ++++++++++++++++++++----------- animatediff/nodes_conditioning.py | 1 + 2 files changed, 59 insertions(+), 33 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 8b8c549..f244364 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -49,10 +49,10 @@ class ModelPatcherAndInjector(ModelPatcher): # lora hook stuff self.hooked_patches = {} # binds LoraHook to specific keys - self.hooked_backup = {} + self.hooked_backup: dict[str, tuple[Tensor, torch.device]] = {} self.cached_hooked_patches = {} # binds LoraHookGroup to pre-calculated weights (speed optimization) self.current_lora_hooks = None - self.lora_hook_mode = LoraHookMode.MAX_SPEED + self.lora_hook_mode = LoraHookMode.MAX_SPEED #LoraHookMode.MIN_VRAM # # replace hook stuff self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys # injection stuff @@ -239,7 +239,6 @@ class ModelPatcherAndInjector(ModelPatcher): logger.warning(f"Cached LoraHook hook could not patch. key doesn't exist in model: {key}") self.patch_cached_hooked_weight(cached_weights=cached_weights, key=key) else: - # TODO: handle lowvram # get combined patches of relevant lora_hooks relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) replace_patches = self.get_hooked_replace_patches(lora_hooks=lora_hooks) @@ -261,15 +260,15 @@ class ModelPatcherAndInjector(ModelPatcher): def patch_cached_hooked_weight(self, cached_weights: dict, key: str): # TODO: handle lowvram + weight: Tensor = comfy.utils.get_attr(self.model, key) + inplace_update = self.weight_inplace_update target_device = self.offload_device if self.lora_hook_mode == LoraHookMode.MAX_SPEED: - target_device = self.current_device - - weight: Tensor = comfy.utils.get_attr(self.model, key) + target_device = weight.device if key not in self.hooked_backup: - self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) + self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) if inplace_update: comfy.utils.copy_to_param(self.model, key, cached_weights[key]) else: @@ -282,16 +281,26 @@ class ModelPatcherAndInjector(ModelPatcher): def patch_hooked_weight_to_device(self, lora_hooks: LoraHookGroup, combined_patches: dict, key: str, device_to=None): if key not in combined_patches: return - + weight: Tensor = comfy.utils.get_attr(self.model, key) + #device_to = weight.device + # if device_to is None: + # device_to = self.hooked_device + # if lowvram and a qualifying key, do lowvram stuff if needed, else continue on with normal behavior + # if self.model_lowvram: + # device_to = weight.device + # if key.endswith(".bias") or key.endswith(".weight"): + # finished = self._lowvram_patch_hooked_weight_to_device(lora_hooks, combined_patches, key, device_to) + # if finished: + # return inplace_update = self.weight_inplace_update target_device = self.offload_device if self.lora_hook_mode == LoraHookMode.MAX_SPEED: - target_device = self.current_device + target_device = weight.device if key not in self.hooked_backup: - self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) + self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) if device_to is not None: temp_weight = comfy.model_management.cast_to_device(weight, device_to, torch.float32, copy=True) @@ -306,6 +315,27 @@ class ModelPatcherAndInjector(ModelPatcher): else: comfy.utils.set_attr_param(self.model, key, out_weight) + def _lowvram_patch_hooked_weight_to_device(self, lora_hooks: LoraHookGroup, combined_patches: dict, key: str, device_to=None): + # get key of module where weight should be stored + split_key = key.split(".") + key_type = split_key[-1] + module_key = ".".join(split_key[:-1]) + module = comfy.utils.get_attr(self.model, module_key) + + if key_type == "bias": + stored_func = getattr(module, "bias_function", None) + if not stored_func: + return False + if module.comfy_cast_weights: + return False + elif key_type == "weight": + stored_func = getattr(module, "weight_function", None) + if module.comfy_cast_weights: + return False + if not stored_func: + return False + return False + def patch_hooked_replace_weight_to_device(self, lora_hooks: LoraHookGroup, model_sd: dict, replace_patches: dict, device_to=None): # first handle replace_patches for key in replace_patches: @@ -316,10 +346,10 @@ class ModelPatcherAndInjector(ModelPatcher): inplace_update = self.weight_inplace_update target_device = self.offload_device if self.lora_hook_mode == LoraHookMode.MAX_SPEED: - target_device = self.current_device + target_device = weight.device if key not in self.hooked_backup: - self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) + self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) out_weight = replace_patches[key].to(self.load_device) if self.lora_hook_mode == LoraHookMode.MAX_SPEED: self.cached_hooked_patches.setdefault(lora_hooks, {}) @@ -339,21 +369,22 @@ class ModelPatcherAndInjector(ModelPatcher): if was_injected: self.eject_model() # TODO: handle lowvram, assuming there is something that needs to be done - if self.model_lowvram: - pass + # if self.model_lowvram: + # logger.warn("lowvram detected in unpatch_hooked!!!") + #if device_to is None: keys = list(self.hooked_backup.keys()) if self.weight_inplace_update: for k in keys: - if device_to is None: - comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) - else: - comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k]) + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: # does not need to be casted - cache device matches needed device + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k][0]) + else: # should be casted as may not match needed device + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1])) else: for k in keys: - if device_to is None: - comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) - else: - comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k]) + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k][0]) + else: # should be casted as may not match needed device + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1])) # clear hooked_backup self.hooked_backup.clear() self.current_lora_hooks = None @@ -412,7 +443,7 @@ class ModelPatcherCLIPHooks(ModelPatcher): # lora hook stuff self.hooked_patches = {} # binds LoraHook to specific keys self.patches_backup = {} - self.hooked_backup = {} + self.hooked_backup: dict[str, tuple[Tensor, torch.device]] = {} self.current_lora_hooks = None self.desired_lora_hooks = None @@ -526,10 +557,10 @@ class ModelPatcherCLIPHooks(ModelPatcher): continue weight: Tensor = comfy.utils.get_attr(self.model, key) inplace_update = self.weight_inplace_update - target_device = self.current_device + target_device = weight.device if key not in self.hooked_backup: - self.hooked_backup[key] = weight.to(device=target_device, copy=inplace_update) + self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) out_weight = replace_patches[key].to(target_device) if inplace_update: comfy.utils.copy_to_param(self.model, key, out_weight) @@ -562,16 +593,10 @@ class ModelPatcherCLIPHooks(ModelPatcher): keys = list(self.hooked_backup.keys()) if self.weight_inplace_update: for k in keys: - if device_to is None: - comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) - else: - comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k]) + comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1])) else: for k in keys: - if device_to is None: - comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k].to(device=self.current_device)) - else: - comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k]) + comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1])) # clear hooked_backup self.hooked_backup.clear() self.current_lora_hooks = None diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 6150fec..6c838e6 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -193,6 +193,7 @@ class ConditioningSetUnmaskedAndCombineHooked: ############################################### + ############################################### ### Scheduling ############################################### From 92b7b8c0d6d8764ea458d3e36c76324410b1e44e Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 26 Apr 2024 17:47:04 -0500 Subject: [PATCH 18/30] Small cleanup of code --- animatediff/model_injection.py | 65 +++++++++++++++------------------- animatediff/utils_motion.py | 4 +-- 2 files changed, 31 insertions(+), 38 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index f244364..4d36f11 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -52,7 +52,7 @@ class ModelPatcherAndInjector(ModelPatcher): self.hooked_backup: dict[str, tuple[Tensor, torch.device]] = {} self.cached_hooked_patches = {} # binds LoraHookGroup to pre-calculated weights (speed optimization) self.current_lora_hooks = None - self.lora_hook_mode = LoraHookMode.MAX_SPEED #LoraHookMode.MIN_VRAM # + self.lora_hook_mode = LoraHookMode.MAX_SPEED # replace hook stuff self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys # injection stuff @@ -254,10 +254,6 @@ class ModelPatcherAndInjector(ModelPatcher): if was_injected: self.inject_model() - def patch_hooked_lowvram(self, lora_hooks: LoraHookGroup, device_to=None, lowvram_model_memory=0): - # TODO: handle lowvram situation - pass - def patch_cached_hooked_weight(self, cached_weights: dict, key: str): # TODO: handle lowvram weight: Tensor = comfy.utils.get_attr(self.model, key) @@ -283,16 +279,10 @@ class ModelPatcherAndInjector(ModelPatcher): return weight: Tensor = comfy.utils.get_attr(self.model, key) - #device_to = weight.device - # if device_to is None: - # device_to = self.hooked_device + # TODO: handle lowvram stuff if necessary # if lowvram and a qualifying key, do lowvram stuff if needed, else continue on with normal behavior # if self.model_lowvram: - # device_to = weight.device - # if key.endswith(".bias") or key.endswith(".weight"): - # finished = self._lowvram_patch_hooked_weight_to_device(lora_hooks, combined_patches, key, device_to) - # if finished: - # return + # self._lowvram_patch_hooked_weight_to_device(lora_hooks, combined_patches, key, device_to) inplace_update = self.weight_inplace_update target_device = self.offload_device @@ -315,27 +305,6 @@ class ModelPatcherAndInjector(ModelPatcher): else: comfy.utils.set_attr_param(self.model, key, out_weight) - def _lowvram_patch_hooked_weight_to_device(self, lora_hooks: LoraHookGroup, combined_patches: dict, key: str, device_to=None): - # get key of module where weight should be stored - split_key = key.split(".") - key_type = split_key[-1] - module_key = ".".join(split_key[:-1]) - module = comfy.utils.get_attr(self.model, module_key) - - if key_type == "bias": - stored_func = getattr(module, "bias_function", None) - if not stored_func: - return False - if module.comfy_cast_weights: - return False - elif key_type == "weight": - stored_func = getattr(module, "weight_function", None) - if module.comfy_cast_weights: - return False - if not stored_func: - return False - return False - def patch_hooked_replace_weight_to_device(self, lora_hooks: LoraHookGroup, model_sd: dict, replace_patches: dict, device_to=None): # first handle replace_patches for key in replace_patches: @@ -370,8 +339,7 @@ class ModelPatcherAndInjector(ModelPatcher): self.eject_model() # TODO: handle lowvram, assuming there is something that needs to be done # if self.model_lowvram: - # logger.warn("lowvram detected in unpatch_hooked!!!") - #if device_to is None: + # self._lowvram_unpatch_hooked() keys = list(self.hooked_backup.keys()) if self.weight_inplace_update: for k in keys: @@ -392,6 +360,31 @@ class ModelPatcherAndInjector(ModelPatcher): if was_injected: self.inject_model() + def _lowvram_patch_hooked_weight_to_device(self, lora_hooks: LoraHookGroup, combined_patches: dict, key: str, device_to=None): + return + # get key of module where weight should be stored + split_key = key.split(".") + key_type = split_key[-1] + module_key = ".".join(split_key[:-1]) + module = comfy.utils.get_attr(self.model, module_key) + + if key_type == "bias": + stored_func = getattr(module, "bias_function", None) + if not stored_func: + return False + if module.comfy_cast_weights: + return False + elif key_type == "weight": + stored_func = getattr(module, "weight_function", None) + if module.comfy_cast_weights: + return False + if not stored_func: + return False + return False + + def _lowvram_unpatch_hooked(self): + return + class CLIPWithHooks(CLIP): def __init__(self, clip: Union[CLIP, 'CLIPWithHooks']): diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index ada4cad..79322ce 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -275,9 +275,9 @@ class ADKeyframeGroup: class LoraHookMode: MIN_VRAM = "min_vram" - MIN_VRAM_LOWVRAM = "min_vram_lowvram" MAX_SPEED = "max_speed" - MAX_SPEED_LOWVRAM = "max_speed_lowvram" + #MIN_VRAM_LOWVRAM = "min_vram_lowvram" + #MAX_SPEED_LOWVRAM = "max_speed_lowvram" class LoraHook: From 2dd285052d4da3a86af3fc275e0aaeecce6cda92 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Fri, 26 Apr 2024 18:03:03 -0500 Subject: [PATCH 19/30] Fixed lowvram I accidentally broke during cleanup --- animatediff/model_injection.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 4d36f11..ccc75e3 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -292,10 +292,10 @@ class ModelPatcherAndInjector(ModelPatcher): if key not in self.hooked_backup: self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) - if device_to is not None: - temp_weight = comfy.model_management.cast_to_device(weight, device_to, torch.float32, copy=True) - else: - temp_weight = weight.to(torch.float32, copy=True) + # if device_to is not None: + temp_weight = comfy.model_management.cast_to_device(weight, weight.device, torch.float32, copy=True) + # else: + # temp_weight = weight.to(torch.float32, copy=True) out_weight = self.calculate_weight(combined_patches[key], temp_weight, key).to(weight.dtype) if self.lora_hook_mode == LoraHookMode.MAX_SPEED: self.cached_hooked_patches.setdefault(lora_hooks, {}) From aba69c0e5560646705c85e56826dcbab03c1e1d0 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 27 Apr 2024 13:14:47 -0500 Subject: [PATCH 20/30] Cleaned up hooked patching code --- animatediff/model_injection.py | 131 +++++++++++++++------------------ 1 file changed, 59 insertions(+), 72 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index ccc75e3..4ce2b8b 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -179,91 +179,80 @@ class ModelPatcherAndInjector(ModelPatcher): else: patched_model = super().patch_model(device_to, patch_weights) # finally, perform motion model injection - self.inject_model(device_to=device_to) + self.inject_model() return patched_model def unpatch_model(self, device_to=None, unpatch_weights=True): # first, eject motion model from unet - self.eject_model(device_to=device_to) + self.eject_model() # finally, do normal model unpatching if unpatch_weights: # TODO: keep only 'else' portion when don't need to worry about past comfy versions # handle hooked_patches first - self.unpatch_hooked(device_to=device_to) + self.unpatch_hooked() self.clear_cached_hooked_weights() return super().unpatch_model(device_to) else: return super().unpatch_model(device_to, unpatch_weights) - def inject_model(self, device_to=None): + def inject_model(self): if self.motion_models is not None: for motion_model in self.motion_models.models: self.currently_injected = True motion_model.model.inject(self) - def eject_model(self, device_to=None): + def eject_model(self): if self.motion_models is not None: for motion_model in self.motion_models.models: motion_model.model.eject(self) self.currently_injected = False - def apply_lora_hooks(self, lora_hooks: LoraHookGroup, device_to=None): + def apply_lora_hooks(self, lora_hooks: LoraHookGroup): # first, determine if need to reapply patches if self.current_lora_hooks == lora_hooks: return - # unpatch hooks, if needed - self.unpatch_hooked(device_to=device_to) - # finally, patch hooks - self.patch_hooked(lora_hooks=lora_hooks, device_to=device_to) + # patch hooks + self.patch_hooked(lora_hooks=lora_hooks) - def patch_hooked(self, lora_hooks: LoraHookGroup, device_to=None, patch_weights=True) -> None: - if not patch_weights: - return - # use current device - if not device_to: - device_to=self.current_device - # then, handle weights - if patch_weights: - # first, unpatch any previous patches - self.unpatch_hooked() - # eject model, if needed - was_injected = self.currently_injected - if was_injected: - self.eject_model() + def patch_hooked(self, lora_hooks: LoraHookGroup) -> None: + # first, unpatch any previous patches + self.unpatch_hooked() + # eject model, if needed + was_injected = self.currently_injected + if was_injected: + self.eject_model() - model_sd = self.model_state_dict() - # if have cached weights for lora_hooks, use it - cached_weights = self.cached_hooked_patches.get(lora_hooks, None) - if cached_weights is not None: - for key in cached_weights: - if key not in model_sd: - logger.warning(f"Cached LoraHook hook could not patch. key doesn't exist in model: {key}") - self.patch_cached_hooked_weight(cached_weights=cached_weights, key=key) - else: - # get combined patches of relevant lora_hooks - relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) - replace_patches = self.get_hooked_replace_patches(lora_hooks=lora_hooks) - if len(replace_patches) > 0: - self.patch_hooked_replace_weight_to_device(lora_hooks=lora_hooks, model_sd=model_sd, replace_patches=replace_patches) - for key in relevant_patches: - if key not in model_sd: - logger.warning(f"LoraHook could not patch. key doesn't exist in model: {key}") - continue - self.patch_hooked_weight_to_device(lora_hooks=lora_hooks, combined_patches=relevant_patches, key=key, device_to=device_to) - self.current_lora_hooks = lora_hooks - # reinject model, if needed - if was_injected: - self.inject_model() + model_sd = self.model_state_dict() + # if have cached weights for lora_hooks, use it + cached_weights = self.cached_hooked_patches.get(lora_hooks, None) + if cached_weights is not None: + for key in cached_weights: + if key not in model_sd: + logger.warning(f"Cached LoraHook could not patch. key doesn't exist in model: {key}") + self.patch_cached_hooked_weight(cached_weights=cached_weights, key=key) + else: + # get combined patches of relevant lora_hooks + relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) + replace_patches = self.get_hooked_replace_patches(lora_hooks=lora_hooks) + if len(replace_patches) > 0: + self.patch_hooked_replace_weight_to_device(lora_hooks=lora_hooks, model_sd=model_sd, replace_patches=replace_patches) + for key in relevant_patches: + if key not in model_sd: + logger.warning(f"LoraHook could not patch. key doesn't exist in model: {key}") + continue + self.patch_hooked_weight_to_device(lora_hooks=lora_hooks, combined_patches=relevant_patches, key=key) + self.current_lora_hooks = lora_hooks + # reinject model, if needed + if was_injected: + self.inject_model() def patch_cached_hooked_weight(self, cached_weights: dict, key: str): # TODO: handle lowvram - weight: Tensor = comfy.utils.get_attr(self.model, key) - inplace_update = self.weight_inplace_update - target_device = self.offload_device - if self.lora_hook_mode == LoraHookMode.MAX_SPEED: - target_device = weight.device - if key not in self.hooked_backup: + weight: Tensor = comfy.utils.get_attr(self.model, key) + target_device = self.offload_device + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + target_device = weight.device self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) if inplace_update: comfy.utils.copy_to_param(self.model, key, cached_weights[key]) @@ -274,24 +263,23 @@ class ModelPatcherAndInjector(ModelPatcher): self.cached_hooked_patches.clear() self.current_lora_hooks = None - def patch_hooked_weight_to_device(self, lora_hooks: LoraHookGroup, combined_patches: dict, key: str, device_to=None): + def patch_hooked_weight_to_device(self, lora_hooks: LoraHookGroup, combined_patches: dict, key: str): if key not in combined_patches: return + inplace_update = self.weight_inplace_update weight: Tensor = comfy.utils.get_attr(self.model, key) + if key not in self.hooked_backup: + target_device = self.offload_device + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + target_device = weight.device + self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) + # TODO: handle lowvram stuff if necessary # if lowvram and a qualifying key, do lowvram stuff if needed, else continue on with normal behavior # if self.model_lowvram: # self._lowvram_patch_hooked_weight_to_device(lora_hooks, combined_patches, key, device_to) - inplace_update = self.weight_inplace_update - target_device = self.offload_device - if self.lora_hook_mode == LoraHookMode.MAX_SPEED: - target_device = weight.device - - if key not in self.hooked_backup: - self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) - # if device_to is not None: temp_weight = comfy.model_management.cast_to_device(weight, weight.device, torch.float32, copy=True) # else: @@ -305,21 +293,22 @@ class ModelPatcherAndInjector(ModelPatcher): else: comfy.utils.set_attr_param(self.model, key, out_weight) - def patch_hooked_replace_weight_to_device(self, lora_hooks: LoraHookGroup, model_sd: dict, replace_patches: dict, device_to=None): + def patch_hooked_replace_weight_to_device(self, lora_hooks: LoraHookGroup, model_sd: dict, replace_patches: dict): # first handle replace_patches for key in replace_patches: if key not in model_sd: logger.warning(f"LoraHook could not replace patch. key doesn't exist in model: {key}") continue - weight: Tensor = comfy.utils.get_attr(self.model, key) + inplace_update = self.weight_inplace_update - target_device = self.offload_device - if self.lora_hook_mode == LoraHookMode.MAX_SPEED: - target_device = weight.device - + weight: Tensor = comfy.utils.get_attr(self.model, key) if key not in self.hooked_backup: + target_device = self.offload_device + if self.lora_hook_mode == LoraHookMode.MAX_SPEED: + target_device = weight.device self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) - out_weight = replace_patches[key].to(self.load_device) + + out_weight = replace_patches[key].to(weight.device) if self.lora_hook_mode == LoraHookMode.MAX_SPEED: self.cached_hooked_patches.setdefault(lora_hooks, {}) self.cached_hooked_patches[lora_hooks][key] = out_weight @@ -328,9 +317,7 @@ class ModelPatcherAndInjector(ModelPatcher): else: comfy.utils.set_attr_param(self.model, key, out_weight) - def unpatch_hooked(self, device_to=None, unpatch_weights=True) -> None: - if not unpatch_weights: - return + def unpatch_hooked(self) -> None: # if no backups from before hook, then nothing to unpatch if len(self.hooked_backup) == 0: return From df81a94c84bcd6eb456f0c5224c93d81be2c675c Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sat, 27 Apr 2024 18:23:02 -0500 Subject: [PATCH 21/30] More cleanup, prepared scaffolding for tracking which modules use lowvram weight_function and bias_function --- animatediff/model_injection.py | 93 +++++++++++++++++++--------------- 1 file changed, 53 insertions(+), 40 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 4ce2b8b..e9f1c28 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -53,6 +53,8 @@ class ModelPatcherAndInjector(ModelPatcher): self.cached_hooked_patches = {} # binds LoraHookGroup to pre-calculated weights (speed optimization) self.current_lora_hooks = None self.lora_hook_mode = LoraHookMode.MAX_SPEED + self.model_params_lowvram = False + self.model_params_lowvram_keys = {} # keeps track of keys with applied 'weight_function' or 'bias_function' # replace hook stuff self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys # injection stuff @@ -182,6 +184,22 @@ class ModelPatcherAndInjector(ModelPatcher): self.inject_model() return patched_model + def patch_model_lowvram(self, *args, **kwargs): + try: + return super().patch_model_lowvram(*args, **kwargs) + finally: + # check if any modules have weight_function or bias_function that is not None + # NOTE: this serves no purpose currently, but I have it here for future reasons + for n, m in self.model.named_modules(): + if not hasattr(m, "comfy_cast_weights"): + continue + if getattr(m, "weight_function", None) is not None: + self.model_params_lowvram = True + self.model_params_lowvram_keys[f"{n}.weight"] = n + if getattr(m, "bias_function", None) is not None: + self.model_params_lowvram = True + self.model_params_lowvram_keys[f"{n}.weight"] = n + def unpatch_model(self, device_to=None, unpatch_weights=True): # first, eject motion model from unet self.eject_model() @@ -190,9 +208,17 @@ class ModelPatcherAndInjector(ModelPatcher): # handle hooked_patches first self.unpatch_hooked() self.clear_cached_hooked_weights() - return super().unpatch_model(device_to) + try: + return super().unpatch_model(device_to) + finally: + self.model_params_lowvram = False + self.model_params_lowvram_keys.clear() else: - return super().unpatch_model(device_to, unpatch_weights) + try: + return super().unpatch_model(device_to, unpatch_weights) + finally: + self.model_params_lowvram = False + self.model_params_lowvram_keys.clear() def inject_model(self): if self.motion_models is not None: @@ -246,7 +272,7 @@ class ModelPatcherAndInjector(ModelPatcher): self.inject_model() def patch_cached_hooked_weight(self, cached_weights: dict, key: str): - # TODO: handle lowvram + # TODO: handle model_params_lowvram stuff if necessary inplace_update = self.weight_inplace_update if key not in self.hooked_backup: weight: Tensor = comfy.utils.get_attr(self.model, key) @@ -275,15 +301,8 @@ class ModelPatcherAndInjector(ModelPatcher): target_device = weight.device self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device) - # TODO: handle lowvram stuff if necessary - # if lowvram and a qualifying key, do lowvram stuff if needed, else continue on with normal behavior - # if self.model_lowvram: - # self._lowvram_patch_hooked_weight_to_device(lora_hooks, combined_patches, key, device_to) - - # if device_to is not None: + # TODO: handle model_params_lowvram stuff if necessary temp_weight = comfy.model_management.cast_to_device(weight, weight.device, torch.float32, copy=True) - # else: - # temp_weight = weight.to(torch.float32, copy=True) out_weight = self.calculate_weight(combined_patches[key], temp_weight, key).to(weight.dtype) if self.lora_hook_mode == LoraHookMode.MAX_SPEED: self.cached_hooked_patches.setdefault(lora_hooks, {}) @@ -303,6 +322,7 @@ class ModelPatcherAndInjector(ModelPatcher): inplace_update = self.weight_inplace_update weight: Tensor = comfy.utils.get_attr(self.model, key) if key not in self.hooked_backup: + # TODO: handle model_params_lowvram stuff if necessary target_device = self.offload_device if self.lora_hook_mode == LoraHookMode.MAX_SPEED: target_device = weight.device @@ -324,9 +344,7 @@ class ModelPatcherAndInjector(ModelPatcher): was_injected = self.currently_injected if was_injected: self.eject_model() - # TODO: handle lowvram, assuming there is something that needs to be done - # if self.model_lowvram: - # self._lowvram_unpatch_hooked() + # TODO: handle model_params_lowvram stuff if necessary keys = list(self.hooked_backup.keys()) if self.weight_inplace_update: for k in keys: @@ -347,31 +365,6 @@ class ModelPatcherAndInjector(ModelPatcher): if was_injected: self.inject_model() - def _lowvram_patch_hooked_weight_to_device(self, lora_hooks: LoraHookGroup, combined_patches: dict, key: str, device_to=None): - return - # get key of module where weight should be stored - split_key = key.split(".") - key_type = split_key[-1] - module_key = ".".join(split_key[:-1]) - module = comfy.utils.get_attr(self.model, module_key) - - if key_type == "bias": - stored_func = getattr(module, "bias_function", None) - if not stored_func: - return False - if module.comfy_cast_weights: - return False - elif key_type == "weight": - stored_func = getattr(module, "weight_function", None) - if module.comfy_cast_weights: - return False - if not stored_func: - return False - return False - - def _lowvram_unpatch_hooked(self): - return - class CLIPWithHooks(CLIP): def __init__(self, clip: Union[CLIP, 'CLIPWithHooks']): @@ -428,6 +421,9 @@ class ModelPatcherCLIPHooks(ModelPatcher): self.current_lora_hooks = None self.desired_lora_hooks = None self.lora_hook_mode = LoraHookMode.MAX_SPEED + + self.model_params_lowvram = False + self.model_params_lowvram_keys = {} # keeps track of keys with applied 'weight_function' or 'bias_function' # replace hook stuff self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys @@ -479,7 +475,6 @@ class ModelPatcherCLIPHooks(ModelPatcher): ''' Based on add_patches, but for hooked weights. ''' - # TODO: make this work with timestep scheduling current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {}) p = set() for key in patches: @@ -563,6 +558,22 @@ class ModelPatcherCLIPHooks(ModelPatcher): self.current_lora_hooks = self.desired_lora_hooks return super().patch_model(device_to, patch_weights, *args, **kwargs) + def patch_model_lowvram(self, *args, **kwargs): + try: + return super().patch_model_lowvram(*args, **kwargs) + finally: + # check if any modules have weight_function or bias_function that is not None + # NOTE: this serves no purpose currently, but I have it here for future reasons + for n, m in self.model.named_modules(): + if not hasattr(m, "comfy_cast_weights"): + continue + if getattr(m, "weight_function", None) is not None: + self.model_params_lowvram = True + self.model_params_lowvram_keys[f"{n}.weight"] = n + if getattr(m, "bias_function", None) is not None: + self.model_params_lowvram = True + self.model_params_lowvram_keys[f"{n}.weight"] = n + def unpatch_model(self, device_to=None, unpatch_weights=True, *args, **kwargs): try: return super().unpatch_model(device_to, unpatch_weights, *args, **kwargs) @@ -577,6 +588,8 @@ class ModelPatcherCLIPHooks(ModelPatcher): else: for k in keys: comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1])) + self.model_params_lowvram = False + self.model_params_lowvram_keys.clear() # clear hooked_backup self.hooked_backup.clear() self.current_lora_hooks = None From 005d49a0b2e7830b4670699f798a7f30bd454189 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 28 Apr 2024 05:06:15 -0500 Subject: [PATCH 22/30] Moved LoraHook stuff into conditioning.py from utilts_motion.py --- animatediff/conditioning.py | 66 ++++++++++++++++++++++++++++++- animatediff/model_injection.py | 4 +- animatediff/nodes_conditioning.py | 4 +- animatediff/sample_settings.py | 3 +- animatediff/sampling.py | 4 +- animatediff/utils_motion.py | 66 ------------------------------- 6 files changed, 73 insertions(+), 74 deletions(-) diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py index 7c0c39b..ab72aeb 100644 --- a/animatediff/conditioning.py +++ b/animatediff/conditioning.py @@ -1,6 +1,70 @@ +import uuid from torch import Tensor -from .utils_motion import LoraHookGroup + +class LoraHookMode: + MIN_VRAM = "min_vram" + MAX_SPEED = "max_speed" + #MIN_VRAM_LOWVRAM = "min_vram_lowvram" + #MAX_SPEED_LOWVRAM = "max_speed_lowvram" + + +class LoraHook: + def __init__(self, lora_name: str): + self.lora_name = lora_name + self.id = f"{lora_name}|{uuid.uuid4()}" + + # def __eq__(self, other: 'LoraHook'): + # return self.id == other.id + + +class LoraHookGroup: + ''' + Stores LoRA hooks to apply for conditioning + ''' + def __init__(self): + self.hooks: list[LoraHook] = [] + + def names(self): + names = [] + for hook in self.hooks: + names.append(hook.lora_name) + return ",".join(names) + + def add(self, hook: str): + if hook not in self.hooks: + self.hooks.append(hook) + + def is_empty(self): + return len(self.hooks) == 0 + + def clone(self): + cloned = LoraHookGroup() + for hook in self.hooks: + cloned.add(hook) + return cloned + + def clone_and_combine(self, other: 'LoraHookGroup'): + cloned = self.clone() + for hook in other.hooks: + cloned.add(hook) + return cloned + + @staticmethod + def combine_all_lora_hooks(lora_hooks_list: list['LoraHookGroup'], require_count=2) -> 'LoraHookGroup': + actual: list[LoraHookGroup] = [] + for group in lora_hooks_list: + if group is not None: + actual.append(group) + if len(actual) < require_count: + raise Exception(f"Need at least {require_count} LoRA Hooks to combine, but only had {len(actual)}.") + final_hook: LoraHookGroup = None + for hook in actual: + if final_hook is None: + final_hook = hook.clone() + else: + final_hook = final_hook.clone_and_combine(hook) + return final_hook class COND_CONST: diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index e9f1c28..9628568 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -20,8 +20,8 @@ from .context import ContextOptions, ContextOptions, ContextOptionsGroup from .motion_module_ad import (AnimateDiffModel, AnimateDiffFormat, EncoderOnlyAnimateDiffModel, VersatileAttention, has_mid_block, normalize_ad_state_dict, get_position_encoding_max_len) from .logger import logger -from .utils_motion import (ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max, - LoraHook, LoraHookGroup, LoraHookMode) +from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max +from .conditioning import LoraHook, LoraHookGroup, LoraHookMode from .motion_lora import MotionLoraInfo, MotionLoraList from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type from .sample_settings import SampleSettings, SeedNoiseGeneration diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 6c838e6..2f9d765 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -8,9 +8,9 @@ from comfy.sd import CLIP import comfy.sd import comfy.utils -from .conditioning import (COND_CONST, TimestepsCond, set_mask_conds, set_mask_and_combine_conds, set_unmasked_and_combine_conds,) +from .conditioning import (COND_CONST, TimestepsCond, set_mask_conds, set_mask_and_combine_conds, set_unmasked_and_combine_conds, + LoraHook, LoraHookGroup) from .model_injection import ModelPatcherAndInjector, CLIPWithHooks, load_hooked_lora_for_models, load_model_as_hooked_lora_for_models -from .utils_motion import LoraHook, LoraHookGroup ############################################### diff --git a/animatediff/sample_settings.py b/animatediff/sample_settings.py index c74c376..f46d907 100644 --- a/animatediff/sample_settings.py +++ b/animatediff/sample_settings.py @@ -9,9 +9,10 @@ from comfy.model_patcher import ModelPatcher from comfy.model_base import BaseModel from . import freeinit +from .conditioning import LoraHookMode from .context import ContextOptions, ContextOptionsGroup from .utils_model import SigmaSchedule -from .utils_motion import extend_to_batch_size, get_sorted_list_via_attr, prepare_mask_batch, LoraHookMode +from .utils_motion import extend_to_batch_size, get_sorted_list_via_attr, prepare_mask_batch from .logger import logger diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 05b3504..a1448be 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -22,11 +22,11 @@ from comfy.controlnet import ControlBase from comfy.model_base import BaseModel import comfy.ops -from .conditioning import COND_CONST +from .conditioning import COND_CONST, LoraHookGroup from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows from .sample_settings import IterationOptions, SampleSettings, SeedNoiseGeneration from .utils_model import ModelTypeSD -from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher, LoraHookGroup +from .model_injection import InjectionParams, ModelPatcherAndInjector, MotionModelGroup, MotionModelPatcher from .motion_module_ad import AnimateDiffFormat, AnimateDiffInfo, AnimateDiffVersion, VanillaTemporalModule from .logger import logger diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index 79322ce..2c341b8 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -2,7 +2,6 @@ from typing import Union import torch import torch.nn.functional as F from torch import Tensor, nn -import uuid import comfy.model_management as model_management import comfy.ops @@ -273,71 +272,6 @@ class ADKeyframeGroup: return cloned -class LoraHookMode: - MIN_VRAM = "min_vram" - MAX_SPEED = "max_speed" - #MIN_VRAM_LOWVRAM = "min_vram_lowvram" - #MAX_SPEED_LOWVRAM = "max_speed_lowvram" - - -class LoraHook: - def __init__(self, lora_name: str): - self.lora_name = lora_name - self.id = f"{lora_name}|{uuid.uuid4()}" - - # def __eq__(self, other: 'LoraHook'): - # return self.id == other.id - - -class LoraHookGroup: - ''' - Stores LoRA hooks to apply for conditioning - ''' - def __init__(self): - self.hooks: list[LoraHook] = [] - - def names(self): - names = [] - for hook in self.hooks: - names.append(hook.lora_name) - return ",".join(names) - - def add(self, hook: str): - if hook not in self.hooks: - self.hooks.append(hook) - - def is_empty(self): - return len(self.hooks) == 0 - - def clone(self): - cloned = LoraHookGroup() - for hook in self.hooks: - cloned.add(hook) - return cloned - - def clone_and_combine(self, other: 'LoraHookGroup'): - cloned = self.clone() - for hook in other.hooks: - cloned.add(hook) - return cloned - - @staticmethod - def combine_all_lora_hooks(lora_hooks_list: list['LoraHookGroup'], require_count=2) -> 'LoraHookGroup': - actual: list[LoraHookGroup] = [] - for group in lora_hooks_list: - if group is not None: - actual.append(group) - if len(actual) < require_count: - raise Exception(f"Need at least {require_count} LoRA Hooks to combine, but only had {len(actual)}.") - final_hook: LoraHookGroup = None - for hook in actual: - if final_hook is None: - final_hook = hook.clone() - else: - final_hook = final_hook.clone_and_combine(hook) - return final_hook - - class DummyNNModule(nn.Module): class DoNothingWhenCalled: def __call__(self, *args, **kwargs): From b1893ca36c2ee63315c77fd1e7ad470b2192abfc Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 28 Apr 2024 06:09:28 -0500 Subject: [PATCH 23/30] Fixed default_conds behavior at CFG=1.0, start of LoRA scheduling work --- animatediff/conditioning.py | 91 +++++++++++++++++++++++++++++++++++++ animatediff/sampling.py | 17 ++++--- animatediff/utils_motion.py | 2 +- 3 files changed, 103 insertions(+), 7 deletions(-) diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py index ab72aeb..b5f8cfa 100644 --- a/animatediff/conditioning.py +++ b/animatediff/conditioning.py @@ -1,6 +1,10 @@ import uuid from torch import Tensor +from comfy.model_base import BaseModel + +from .utils_motion import get_sorted_list_via_attr + class LoraHookMode: MIN_VRAM = "min_vram" @@ -67,6 +71,93 @@ class LoraHookGroup: return final_hook +class LoraHookKeyframe: + def __init__(self, strength: float, start_percent=0.0, guarantee_steps=1): + self.strength = strength + # scheduling + self.start_percent = float(start_percent) + self.start_t = 999999999.9 + self.guarantee_steps = guarantee_steps + + def clone(self): + c = LoraHookKeyframe(strength=self.strength, + start_percent=self.start_percent, guarantee_steps=self.guarantee_steps) + c.start_t = self.start_t + return c + +class LoraHookKeyframeGroup: + def __init__(self): + self.keyframes: list[LoraHookKeyframe] = [] + self._current_keyframe: LoraHookKeyframe = None + self._current_used_steps: int = 0 + self._current_index: int = 0 + + def reset(self): + self._current_keyframe = None + self._current_used_steps = 0 + self._current_index = 0 + self._set_first_as_current() + + def add(self, keyframe: LoraHookKeyframe): + # add to end of list, then sort + self.keyframes.append(keyframe) + self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent") + self._set_first_as_current() + + def _set_first_as_current(self): + if len(self.keyframes) > 0: + self._current_keyframe = self.keyframes[0] + else: + self._current_keyframe = None + + def has_index(self, index: int) -> int: + return index >= 0 and index < len(self.keyframes) + + def is_empty(self) -> bool: + return len(self.keyframes) == 0 + + def clone(self): + cloned = LoraHookKeyframeGroup() + for keyframe in self.keyframes: + cloned.keyframes.append(keyframe) + cloned._set_first_as_current() + return cloned + + def initialize_timesteps(self, model: BaseModel): + for keyframe in self.keyframes: + keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent) + + def prepare_current_keyframe(self, t: Tensor): + curr_t: float = t[0] + prev_index = self._current_index + # if met guaranteed steps, look for next keyframe in case need to switch + if self._current_used_steps >= self._current_keyframe.guarantee_steps: + # if has next index, loop through and see if need t oswitch + if self.has_index(self._current_index+1): + for i in range(self._current_index+1, len(self.keyframes)): + eval_c = self.keyframes[i] + # check if start_t is greater or equal to curr_t + # NOTE: t is in terms of sigmas, not percent, so bigger number = earlier step in sampling + if eval_c.start_t >= curr_t: + self._current_index = i + self._current_keyframe = eval_c + self._current_used_steps = 0 + # if guarantee_steps greater than zero, stop searching for other keyframes + if self._current_keyframe.guarantee_steps > 0: + break + # if eval_c is outside the percent range, stop looking further + else: break + # update steps current context is used + self._current_used_steps += 1 + + # properties shadow those of LoraHookKeyframe + @property + def strength(self): + if self._current_keyframe is not None: + return self._current_keyframe.strength + return None + + class COND_CONST: KEY_LORA_HOOK = "lora_hook" KEY_DEFAULT_COND = "default_cond" diff --git a/animatediff/sampling.py b/animatediff/sampling.py index a1448be..4af7771 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -714,12 +714,17 @@ def separate_default_conds(conds: list[dict]): for i in range(len(conds)): c = [] default_c = [] - for t in conds[i]: - # check if cond is a default cond - if COND_CONST.KEY_DEFAULT_COND in t: - default_c.append(t) - else: - c.append(t) + # if cond is None, make normal/default_conds reflect that too + if conds[i] is None: + c = None + default_c = [] + else: + for t in conds[i]: + # check if cond is a default cond + if COND_CONST.KEY_DEFAULT_COND in t: + default_c.append(t) + else: + c.append(t) normal_conds.append(c) default_conds.append(default_c) return normal_conds, default_conds diff --git a/animatediff/utils_motion.py b/animatediff/utils_motion.py index 2c341b8..a2d7510 100644 --- a/animatediff/utils_motion.py +++ b/animatediff/utils_motion.py @@ -157,7 +157,7 @@ def get_sorted_list_via_attr(objects: list, attr: str) -> list: unique_attrs = {} for o in objects: val_attr = getattr(o, attr) - attr_list = unique_attrs.get(val_attr, list()) + attr_list: list = unique_attrs.get(val_attr, list()) attr_list.append(o) if val_attr not in unique_attrs: unique_attrs[val_attr] = attr_list From 33d8a98d77b568ec7b40b14ed97ad018ba108e1f Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 28 Apr 2024 10:43:01 -0500 Subject: [PATCH 24/30] Refactored Registered Model As LoRA Hook to work as a normal lora via weight diff instead of hooked_replace, completely removed hooked_replace code --- animatediff/model_injection.py | 131 ++++++++++++------------------ animatediff/nodes_conditioning.py | 13 ++- 2 files changed, 59 insertions(+), 85 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 9628568..c80703b 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -55,8 +55,6 @@ class ModelPatcherAndInjector(ModelPatcher): self.lora_hook_mode = LoraHookMode.MAX_SPEED self.model_params_lowvram = False self.model_params_lowvram_keys = {} # keeps track of keys with applied 'weight_function' or 'bias_function' - # replace hook stuff - self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys # injection stuff self.currently_injected = False self.motion_injection_params: InjectionParams = InjectionParams() @@ -75,11 +73,6 @@ class ModelPatcherAndInjector(ModelPatcher): cloned.cached_hooked_patches[group] = {} for k in self.cached_hooked_patches[group]: cloned.cached_hooked_patches[group][k] = self.cached_hooked_patches[group][k] - # copy replace lora hooks - for hook in self.hooked_replace_patches: - cloned.hooked_replace_patches[hook] = {} - for k in self.hooked_replace_patches[hook]: - cloned.hooked_replace_patches[hook][k] = self.hooked_replace_patches[hook][k] cloned.hooked_backup = self.hooked_backup cloned.current_lora_hooks = self.current_lora_hooks cloned.currently_injected = self.currently_injected @@ -102,7 +95,7 @@ class ModelPatcherAndInjector(ModelPatcher): if not returned: return returned # currently, hook patches require that model gets loaded when sampled, so always say is not a clone if hooks present - if len(self.hooked_patches) > 0 or len(self.hooked_replace_patches) > 0: + if len(self.hooked_patches) > 0: return False if type(self) != type(clone): return False @@ -110,8 +103,6 @@ class ModelPatcherAndInjector(ModelPatcher): return False if self.hooked_patches.keys() != clone.hooked_patches.keys(): return False - if self.hooked_replace_patches.keys() != clone.hooked_replace_patches.keys(): - return False return returned def set_lora_hook_mode(self, lora_hook_mode: str): @@ -135,6 +126,26 @@ class ModelPatcherAndInjector(ModelPatcher): self.patches_uuid = uuid.uuid4() return list(p) + def add_hooked_patches_as_diffs(self, lora_hook: LoraHook, patches: dict, strength_patch=1.0, strength_model=1.0): + ''' + Based on add_hooked_patches, but intended for using a model's weights as lora hook. + ''' + # TODO: make this work with timestep scheduling + current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {}) + p = set() + for key in patches: + if key in self.model_keys: + p.add(key) + current_patches: list[tuple] = current_hooked_patches.get(key, []) + # take difference between desired weight and existing weight to get diff + + current_patches.append((strength_patch, (patches[key]-comfy.utils.get_attr(self.model, key),), strength_model)) + current_hooked_patches[key] = current_patches + self.hooked_patches[lora_hook] = current_hooked_patches + # since should care about these patches too to determine if same model, reroll patches_uuid + self.patches_uuid = uuid.uuid4() + return list(p) + def get_combined_hooked_patches(self, lora_hooks: LoraHookGroup): ''' Returns patches for selected lora_hooks. @@ -150,27 +161,6 @@ class ModelPatcherAndInjector(ModelPatcher): combined_patches[key] = current_patches return combined_patches - def add_hooked_replace_patches(self, lora_hook: LoraHook, patches: dict): - self.hooked_replace_patches.setdefault(lora_hook, {}) - p = set() - for key in patches: - if key in self.model_keys: - p.add(key) - self.hooked_replace_patches[lora_hook][key] = patches[key] - # since should care about these patches too to determine if same model, reroll patches_uuid - self.patches_uuid = uuid.uuid4() - return list(p) - - def get_hooked_replace_patches(self, lora_hooks: LoraHookGroup): - # return first hook found in hooked_replace_patches - patches = {} - if lora_hooks is not None: - for hook in lora_hooks.hooks: - if hook in self.hooked_replace_patches: - patches = self.hooked_replace_patches[hook] - break - return patches - def model_patches_to(self, device): super().model_patches_to(device) @@ -258,9 +248,6 @@ class ModelPatcherAndInjector(ModelPatcher): else: # get combined patches of relevant lora_hooks relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks) - replace_patches = self.get_hooked_replace_patches(lora_hooks=lora_hooks) - if len(replace_patches) > 0: - self.patch_hooked_replace_weight_to_device(lora_hooks=lora_hooks, model_sd=model_sd, replace_patches=replace_patches) for key in relevant_patches: if key not in model_sd: logger.warning(f"LoraHook could not patch. key doesn't exist in model: {key}") @@ -387,13 +374,9 @@ class CLIPWithHooks(CLIP): 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 add_hooked_replace_patches(self, lora_hook: LoraHook, patches): - return self.patcher.add_hooked_replace_patches(lora_hook=lora_hook, patches=patches) - - # def encode_from_tokens(self, tokens, return_pooled=False): - # comfy.model_management.cleanup_models() - # return super().encode_from_tokens(tokens, return_pooled) + + def add_hooked_patches_as_diffs(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): + return self.patcher.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=patches, strength_patch=strength_patch, strength_model=strength_model) class ModelPatcherCLIPHooks(ModelPatcher): @@ -424,8 +407,6 @@ class ModelPatcherCLIPHooks(ModelPatcher): self.model_params_lowvram = False self.model_params_lowvram_keys = {} # keeps track of keys with applied 'weight_function' or 'bias_function' - # replace hook stuff - self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys def clone(self): cloned = ModelPatcherCLIPHooks(self) @@ -434,11 +415,6 @@ class ModelPatcherCLIPHooks(ModelPatcher): cloned.hooked_patches[hook] = {} for k in self.hooked_patches[hook]: cloned.hooked_patches[hook][k] = self.hooked_patches[hook][k][:] - # copy replace lora hooks - for hook in self.hooked_replace_patches: - cloned.hooked_replace_patches[hook] = {} - for k in self.hooked_replace_patches[hook]: - cloned.hooked_replace_patches[hook][k] = self.hooked_replace_patches[hook][k] cloned.patches_backup = self.patches_backup cloned.hooked_backup = self.hooked_backup cloned.current_lora_hooks = self.current_lora_hooks @@ -464,8 +440,6 @@ class ModelPatcherCLIPHooks(ModelPatcher): return False if self.hooked_patches.keys() != clone.hooked_patches.keys(): return False - if self.hooked_replace_patches.keys() != clone.hooked_replace_patches.keys(): - return False return returned def set_desired_hooks(self, lora_hooks: LoraHookGroup): @@ -488,6 +462,24 @@ class ModelPatcherCLIPHooks(ModelPatcher): self.patches_uuid = uuid.uuid4() return list(p) + def add_hooked_patches_as_diffs(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): + ''' + Based on add_hooked_patches, but intended for using a model's weights as lora hook. + ''' + current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {}) + p = set() + for key in patches: + if key in self.model_keys: + p.add(key) + current_patches: list[tuple] = current_hooked_patches.get(key, []) + # take difference between desired weight and existing weight to get diff + current_patches.append((strength_patch, (patches[key]-comfy.utils.get_attr(self.model, key),), strength_model)) + current_hooked_patches[key] = current_patches + self.hooked_patches[lora_hook] = current_hooked_patches + # since should care about these patches too to determine if same model, reroll patches_uuid + self.patches_uuid = uuid.uuid4() + return list(p) + def get_combined_hooked_patches(self, lora_hooks: LoraHookGroup): ''' Returns patches for selected lora_hooks. @@ -502,27 +494,6 @@ class ModelPatcherCLIPHooks(ModelPatcher): current_patches.extend(hook_patches[key]) combined_patches[key] = current_patches return combined_patches - - def add_hooked_replace_patches(self, lora_hook: LoraHook, patches: dict): - self.hooked_replace_patches.setdefault(lora_hook, {}) - p = set() - for key in patches: - if key in self.model_keys: - p.add(key) - self.hooked_replace_patches[lora_hook][key] = patches[key] - # since should care about these patches too to determine if same model, reroll patches_uuid - self.patches_uuid = uuid.uuid4() - return list(p) - - def get_hooked_replace_patches(self, lora_hooks: LoraHookGroup): - # return first hook found in hooked_replace_patches - patches = {} - if lora_hooks is not None: - for hook in lora_hooks.hooks: - if hook in self.hooked_replace_patches: - patches = self.hooked_replace_patches[hook] - break - return patches def patch_hooked_replace_weight_to_device(self, model_sd: dict, replace_patches: dict): # first handle replace_patches @@ -545,12 +516,6 @@ class ModelPatcherCLIPHooks(ModelPatcher): def patch_model(self, device_to=None, patch_weights=True, *args, **kwargs): if self.desired_lora_hooks is not None: self.patches_backup = self.patches.copy() - # first, handle replace patches # TODO: make work properly for CLIP - replace_patches = self.get_hooked_replace_patches(lora_hooks=self.desired_lora_hooks) - if len(replace_patches) > 0: - model_sd = self.model_state_dict() - self.patch_hooked_replace_weight_to_device(model_sd=model_sd, replace_patches=replace_patches) - # then, handle usual patches relevant_patches = self.get_combined_hooked_patches(lora_hooks=self.desired_lora_hooks) for key in relevant_patches: self.patches.setdefault(key, []) @@ -595,7 +560,8 @@ class ModelPatcherCLIPHooks(ModelPatcher): self.current_lora_hooks = None -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): +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: key_map = comfy.lora.model_lora_keys_unet(model.model, key_map) @@ -624,17 +590,20 @@ def load_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInject return (new_modelpatcher, new_clip) -def load_model_as_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInjector], clip: CLIP, model_loaded: ModelPatcher, clip_loaded: CLIP, lora_hook: LoraHook): +def load_model_as_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcherAndInjector], clip: CLIP, model_loaded: ModelPatcher, clip_loaded: CLIP, lora_hook: LoraHook, + strength_model: float, strength_clip: float): if model is not None and model_loaded is not None: new_modelpatcher = ModelPatcherAndInjector.create_from(model) - k = new_modelpatcher.add_hooked_replace_patches(lora_hook=lora_hook, patches=model_loaded.model.state_dict()) + comfy.model_management.unload_model_clones(new_modelpatcher) + k = new_modelpatcher.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=model_loaded.model.state_dict(), strength_patch=strength_model) else: k = () new_modelpatcher = None if clip is not None and clip_loaded is not None: new_clip = CLIPWithHooks(clip) - k1 = new_clip.add_hooked_replace_patches(lora_hook=lora_hook, patches=clip.cond_stage_model.state_dict()) + comfy.model_management.unload_model_clones(new_clip) + k1 = new_clip.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=clip.cond_stage_model.state_dict(), strength_patch=strength_clip) else: k1 = () new_clip = None diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 2f9d765..15b986f 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -298,6 +298,8 @@ class MaskableSDModelLoader: "model": ("MODEL",), "clip": ("CLIP",), "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ), + "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}), } } @@ -305,7 +307,7 @@ class MaskableSDModelLoader: CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/register lora hooks" FUNCTION = "load_model_as_lora" - def load_model_as_lora(self, model: ModelPatcher, clip: CLIP, ckpt_name: str): + def load_model_as_lora(self, model: ModelPatcher, clip: CLIP, ckpt_name: str, strength_model: float, strength_clip: float): ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) out = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings")) model_loaded = out[0] @@ -316,7 +318,8 @@ class MaskableSDModelLoader: lora_hook_group.add(lora_hook) model_lora, clip_lora = load_model_as_hooked_lora_for_models(model=model, clip=clip, model_loaded=model_loaded, clip_loaded=clip_loaded, - lora_hook=lora_hook) + lora_hook=lora_hook, + strength_model=strength_model, strength_clip=strength_clip) return (model_lora, clip_lora, lora_hook_group) @@ -327,6 +330,7 @@ class MaskableSDModelLoaderModelOnly(MaskableSDModelLoader): "required": { "model": ("MODEL",), "ckpt_name": (folder_paths.get_filename_list("checkpoints"), ), + "strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), } } @@ -334,8 +338,9 @@ class MaskableSDModelLoaderModelOnly(MaskableSDModelLoader): CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/register lora hooks" FUNCTION = "load_model_as_lora_model_only" - def load_model_as_lora_model_only(self, model: ModelPatcher, ckpt_name: str): - model_lora, clip_lora, lora_hook = self.load_model_as_lora(model=model, clip=None, ckpt_name=ckpt_name) + def load_model_as_lora_model_only(self, model: ModelPatcher, ckpt_name: str, strength_model: float): + model_lora, clip_lora, lora_hook = self.load_model_as_lora(model=model, clip=None, ckpt_name=ckpt_name, + strength_model=strength_model, strength_clip=0) return (model_lora, lora_hook) ############################################### ############################################### From 78880ad398388341eb79deab9154376d3b219d19 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 28 Apr 2024 11:08:29 -0500 Subject: [PATCH 25/30] Fixed Model as Hooked LoRA from allowing model_sampling keys to be patches --- animatediff/model_injection.py | 17 +++++++++++++---- 1 file changed, 13 insertions(+), 4 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index c80703b..41457b5 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -595,7 +595,14 @@ def load_model_as_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcher if model is not None and model_loaded is not None: new_modelpatcher = ModelPatcherAndInjector.create_from(model) comfy.model_management.unload_model_clones(new_modelpatcher) - k = new_modelpatcher.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=model_loaded.model.state_dict(), strength_patch=strength_model) + expected_model_keys = model_loaded.model_keys.copy() + patches_model: dict[str, Tensor] = model_loaded.model.state_dict() + # do not include ANY model_sampling components of the model that should act as a patch + for key in list(patches_model.keys()): + if key.startswith("model_sampling"): + expected_model_keys.discard(key) + patches_model.pop(key, None) + k = new_modelpatcher.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=patches_model, strength_patch=strength_model) else: k = () new_modelpatcher = None @@ -603,7 +610,9 @@ def load_model_as_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcher if clip is not None and clip_loaded is not None: new_clip = CLIPWithHooks(clip) comfy.model_management.unload_model_clones(new_clip) - k1 = new_clip.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=clip.cond_stage_model.state_dict(), strength_patch=strength_clip) + expected_clip_keys = clip_loaded.patcher.model_keys.copy() + patches_clip: dict[str, Tensor] = clip.cond_stage_model.state_dict() + k1 = new_clip.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=patches_clip, strength_patch=strength_clip) else: k1 = () new_clip = None @@ -611,11 +620,11 @@ def load_model_as_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcher k = set(k) k1 = set(k1) if model is not None and model_loaded is not None: - for key in model_loaded.model_keys: + for key in expected_model_keys: if key not in k: logger.warning(f"MODEL-AS-LORA NOT LOADED {key}") if clip is not None and clip_loaded is not None: - for key in clip_loaded.patcher.model_keys: + for key in expected_clip_keys: if key not in k1: logger.warning(f"CLIP-AS-LORA NOT LOADED {key}") From 5da9bfda9231bab030aaf273b0fd0981ec54157b Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Sun, 28 Apr 2024 11:56:42 -0500 Subject: [PATCH 26/30] Register Model as LoRA Hook (with clip, not the Model Only version) now works properly thanks to the previous refactor of Model as LoRA Hook code --- animatediff/model_injection.py | 4 ++-- animatediff/nodes.py | 4 ++-- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 41457b5..42194de 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -609,9 +609,9 @@ def load_model_as_hooked_lora_for_models(model: Union[ModelPatcher, ModelPatcher if clip is not None and clip_loaded is not None: new_clip = CLIPWithHooks(clip) - comfy.model_management.unload_model_clones(new_clip) + comfy.model_management.unload_model_clones(new_clip.patcher) expected_clip_keys = clip_loaded.patcher.model_keys.copy() - patches_clip: dict[str, Tensor] = clip.cond_stage_model.state_dict() + patches_clip: dict[str, Tensor] = clip_loaded.cond_stage_model.state_dict() k1 = new_clip.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=patches_clip, strength_patch=strength_clip) else: k1 = () diff --git a/animatediff/nodes.py b/animatediff/nodes.py index 027345c..bc5b89e 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -62,7 +62,7 @@ NODE_CLASS_MAPPINGS = { # Conditioning "ADE_RegisterLoraHook": MaskableLoraLoader, "ADE_RegisterLoraHookModelOnly": MaskableLoraLoaderModelOnly, - #"ADE_RegisterModelAsLoraHook": MaskableSDModelLoader, # CLIP does not work properly on first run + "ADE_RegisterModelAsLoraHook": MaskableSDModelLoader, "ADE_RegisterModelAsLoraHookModelOnly": MaskableSDModelLoaderModelOnly, "ADE_CombineLoraHooks": CombineLoraHooks, "ADE_CombineLoraHooksFour": CombineLoraHookFourOptional, @@ -163,7 +163,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Conditioning "ADE_RegisterLoraHook": "Register LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_RegisterLoraHookModelOnly": "Register LoRA Hook (Model Only) πŸŽ­πŸ…πŸ…“", - #"ADE_RegisterModelAsLoraHook": "Register Model as LoRA HookπŸ”¬ πŸŽ­πŸ…πŸ…“", # CLIP does not work properly on first run + "ADE_RegisterModelAsLoraHook": "Register Model as LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_RegisterModelAsLoraHookModelOnly": "Register Model as LoRA Hook (MO) πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooks": "Combine LoRA Hooks [2] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksFour": "Combine LoRA Hooks [4] πŸŽ­πŸ…πŸ…“", From 910475c7b2f3e8c481773c4c85ada44815d945ab Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 29 Apr 2024 11:01:25 -0500 Subject: [PATCH 27/30] Added LoRA Hook Keyframe support for scheduling any LoRA Hook, whether it be individual or group --- animatediff/conditioning.py | 65 +++++++++++-- animatediff/model_injection.py | 58 ++++++++---- animatediff/nodes.py | 15 ++- animatediff/nodes_conditioning.py | 149 +++++++++++++++++++++++++++++- animatediff/sampling.py | 51 ++++++---- 5 files changed, 291 insertions(+), 47 deletions(-) diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py index b5f8cfa..78dd23f 100644 --- a/animatediff/conditioning.py +++ b/animatediff/conditioning.py @@ -1,4 +1,3 @@ -import uuid from torch import Tensor from comfy.model_base import BaseModel @@ -13,13 +12,42 @@ class LoraHookMode: #MAX_SPEED_LOWVRAM = "max_speed_lowvram" +# Acts simply as a way to track unique LoraHooks +class HookRef: + pass + + class LoraHook: def __init__(self, lora_name: str): self.lora_name = lora_name - self.id = f"{lora_name}|{uuid.uuid4()}" + self.lora_keyframe = LoraHookKeyframeGroup() + self.hook_ref = HookRef() - # def __eq__(self, other: 'LoraHook'): - # return self.id == other.id + def initialize_timesteps(self, model: BaseModel): + self.lora_keyframe.initialize_timesteps(model) + + def reset(self): + self.lora_keyframe.reset() + + + def get_copy(self): + ''' + Copies LoraHook, but maintains same HookRef + ''' + c = LoraHook(lora_name=self.lora_name) + c.lora_keyframe = self.lora_keyframe + c.hook_ref = self.hook_ref # same instance that acts as ref + return c + + @property + def strength(self): + return self.lora_keyframe.strength + + def __eq__(self, other: 'LoraHook'): + return self.__class__ == other.__class__ and self.hook_ref == other.hook_ref + + def __hash__(self): + return hash(self.hook_ref) class LoraHookGroup: @@ -35,24 +63,32 @@ class LoraHookGroup: names.append(hook.lora_name) return ",".join(names) - def add(self, hook: str): + def add(self, hook: LoraHook): if hook not in self.hooks: self.hooks.append(hook) def is_empty(self): return len(self.hooks) == 0 + def contains(self, lora_hook: LoraHook): + return lora_hook in self.hooks + def clone(self): cloned = LoraHookGroup() for hook in self.hooks: - cloned.add(hook) + cloned.add(hook.get_copy()) return cloned def clone_and_combine(self, other: 'LoraHookGroup'): cloned = self.clone() for hook in other.hooks: - cloned.add(hook) + cloned.add(hook.get_copy()) return cloned + + def set_keyframes_on_hooks(self, hook_kf: 'LoraHookKeyframeGroup'): + hook_kf = hook_kf.clone() + for hook in self.hooks: + hook.lora_keyframe = hook_kf @staticmethod def combine_all_lora_hooks(lora_hooks_list: list['LoraHookGroup'], require_count=2) -> 'LoraHookGroup': @@ -91,11 +127,13 @@ class LoraHookKeyframeGroup: self._current_keyframe: LoraHookKeyframe = None self._current_used_steps: int = 0 self._current_index: int = 0 + self._curr_t: float = -1 def reset(self): self._current_keyframe = None self._current_used_steps = 0 self._current_index = 0 + self._curr_t = -1 self._set_first_as_current() def add(self, keyframe: LoraHookKeyframe): @@ -127,8 +165,11 @@ class LoraHookKeyframeGroup: for keyframe in self.keyframes: keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent) - def prepare_current_keyframe(self, t: Tensor): - curr_t: float = t[0] + def prepare_current_keyframe(self, curr_t: float) -> bool: + if self.is_empty(): + return False + if curr_t == self._curr_t: + return False prev_index = self._current_index # if met guaranteed steps, look for next keyframe in case need to switch if self._current_used_steps >= self._current_keyframe.guarantee_steps: @@ -149,13 +190,17 @@ class LoraHookKeyframeGroup: else: break # update steps current context is used self._current_used_steps += 1 + # update current timestep this was performed on + self._curr_t = curr_t + # return True if keyframe changed, False if no change + return prev_index != self._current_index # properties shadow those of LoraHookKeyframe @property def strength(self): if self._current_keyframe is not None: return self._current_keyframe.strength - return None + return 1.0 class COND_CONST: diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index 42194de..0f3d905 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -6,6 +6,7 @@ from torch import Tensor import torch.nn.functional as F import torch import uuid +import math import comfy.lora import comfy.model_management @@ -21,7 +22,7 @@ from .motion_module_ad import (AnimateDiffModel, AnimateDiffFormat, EncoderOnlyA has_mid_block, normalize_ad_state_dict, get_position_encoding_max_len) from .logger import logger from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max -from .conditioning import LoraHook, LoraHookGroup, LoraHookMode +from .conditioning import HookRef, LoraHook, LoraHookGroup, LoraHookMode from .motion_lora import MotionLoraInfo, MotionLoraList from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type from .sample_settings import SampleSettings, SeedNoiseGeneration @@ -48,9 +49,9 @@ class ModelPatcherAndInjector(ModelPatcher): self.object_patches_backup = m.object_patches_backup # lora hook stuff - self.hooked_patches = {} # binds LoraHook to specific keys + self.hooked_patches: dict[HookRef] = {} # binds LoraHook to specific keys self.hooked_backup: dict[str, tuple[Tensor, torch.device]] = {} - self.cached_hooked_patches = {} # binds LoraHookGroup to pre-calculated weights (speed optimization) + self.cached_hooked_patches: dict[LoraHookGroup, dict[str, Tensor]] = {} # binds LoraHookGroup to pre-calculated weights (speed optimization) self.current_lora_hooks = None self.lora_hook_mode = LoraHookMode.MAX_SPEED self.model_params_lowvram = False @@ -64,10 +65,10 @@ class ModelPatcherAndInjector(ModelPatcher): def clone(self, hooks_only=False): cloned = ModelPatcherAndInjector(self) # copy lora hooks - for hook in self.hooked_patches: - cloned.hooked_patches[hook] = {} - for k in self.hooked_patches[hook]: - cloned.hooked_patches[hook][k] = self.hooked_patches[hook][k][:] + for hook_ref in self.hooked_patches: + cloned.hooked_patches[hook_ref] = {} + for k in self.hooked_patches[hook_ref]: + cloned.hooked_patches[hook_ref][k] = self.hooked_patches[hook_ref][k][:] # copy pre-calc weights bound to LoraHookGroups for group in self.cached_hooked_patches: cloned.cached_hooked_patches[group] = {} @@ -107,13 +108,31 @@ class ModelPatcherAndInjector(ModelPatcher): def set_lora_hook_mode(self, lora_hook_mode: str): self.lora_hook_mode = lora_hook_mode + + def prepare_hooked_patches_current_keyframe(self, t: Tensor, hook_groups: list[LoraHookGroup]): + curr_t = t[0] + for hook_group in hook_groups: + for hook in hook_group.hooks: + changed = hook.lora_keyframe.prepare_current_keyframe(curr_t=curr_t) + # if keyframe changed, remove any cached LoraHookGroups that contain hook with the same hook_ref; + # this will cause the weights to be recalculated when sampling + if changed: + for cached_group in list(self.cached_hooked_patches.keys()): + if cached_group.contains(hook): + self.cached_hooked_patches.pop(cached_group) + + def clean_hooks(self): + self.unpatch_hooked() + self.clear_cached_hooked_weights() + # for lora_hook in self.hooked_patches: + # lora_hook.reset() def add_hooked_patches(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0): ''' Based on add_patches, but for hooked weights. ''' # TODO: make this work with timestep scheduling - current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {}) + current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook.hook_ref, {}) p = set() for key in patches: if key in self.model_keys: @@ -121,7 +140,7 @@ class ModelPatcherAndInjector(ModelPatcher): current_patches: list[tuple] = current_hooked_patches.get(key, []) current_patches.append((strength_patch, patches[key], strength_model)) current_hooked_patches[key] = current_patches - self.hooked_patches[lora_hook] = current_hooked_patches + self.hooked_patches[lora_hook.hook_ref] = current_hooked_patches # since should care about these patches too to determine if same model, reroll patches_uuid self.patches_uuid = uuid.uuid4() return list(p) @@ -131,17 +150,16 @@ class ModelPatcherAndInjector(ModelPatcher): Based on add_hooked_patches, but intended for using a model's weights as lora hook. ''' # TODO: make this work with timestep scheduling - current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {}) + current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook.hook_ref, {}) p = set() for key in patches: if key in self.model_keys: p.add(key) current_patches: list[tuple] = current_hooked_patches.get(key, []) # take difference between desired weight and existing weight to get diff - current_patches.append((strength_patch, (patches[key]-comfy.utils.get_attr(self.model, key),), strength_model)) current_hooked_patches[key] = current_patches - self.hooked_patches[lora_hook] = current_hooked_patches + self.hooked_patches[lora_hook.hook_ref] = current_hooked_patches # since should care about these patches too to determine if same model, reroll patches_uuid self.patches_uuid = uuid.uuid4() return list(p) @@ -154,10 +172,19 @@ class ModelPatcherAndInjector(ModelPatcher): combined_patches = {} if lora_hooks is not None: for hook in lora_hooks.hooks: - hook_patches: dict = self.hooked_patches.get(hook, {}) + hook_patches: dict = self.hooked_patches.get(hook.hook_ref, {}) for key in hook_patches.keys(): current_patches: list[tuple] = combined_patches.get(key, []) - current_patches.extend(hook_patches[key]) + if math.isclose(hook.strength, 1.0): + # if hook strength is 1.0, can just add it directly + current_patches.extend(hook_patches[key]) + else: + # otherwise, need to multiply original patch strength by hook strength + # patches are stored as tuples: (strength_patch, (tuple_with_weights,), strength_model) + for patch in hook_patches[key]: + new_patch = list(patch) + new_patch[0] *= hook.strength + current_patches.append(tuple(new_patch)) combined_patches[key] = current_patches return combined_patches @@ -196,8 +223,7 @@ class ModelPatcherAndInjector(ModelPatcher): # finally, do normal model unpatching if unpatch_weights: # TODO: keep only 'else' portion when don't need to worry about past comfy versions # handle hooked_patches first - self.unpatch_hooked() - self.clear_cached_hooked_weights() + self.clean_hooks() try: return super().unpatch_model(device_to) finally: diff --git a/animatediff/nodes.py b/animatediff/nodes.py index bc5b89e..ada1f0f 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -16,7 +16,8 @@ from .nodes_conditioning import (MaskableLoraLoader, MaskableLoraLoaderModelOnly PairedConditioningSetMaskHooked, ConditioningSetMaskHooked, PairedConditioningSetMaskAndCombineHooked, ConditioningSetMaskAndCombineHooked, PairedConditioningSetUnmaskedAndCombineHooked, ConditioningSetUnmaskedAndCombineHooked, - ConditioningTimestepsNode) + ConditioningTimestepsNode, SetLoraHookKeyframes, + CreateLoraHookKeyframe, CreateLoraHookKeyframeInterpolation, CreateLoraHookKeyframeFromStrengthList) from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode, CustomCFGNode, CustomCFGKeyframeNode) from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode) @@ -67,8 +68,12 @@ NODE_CLASS_MAPPINGS = { "ADE_CombineLoraHooks": CombineLoraHooks, "ADE_CombineLoraHooksFour": CombineLoraHookFourOptional, "ADE_CombineLoraHooksEight": CombineLoraHookEightOptional, - "ADE_AttachLoraHookToConditioning": SetModelLoraHook, "ADE_AttachLoraHookToCLIP": SetClipLoraHook, + "ADE_LoraHookKeyframe": CreateLoraHookKeyframe, + "ADE_LoraHookKeyframeInterpolation": CreateLoraHookKeyframeInterpolation, + "ADE_LoraHookKeyframeFromStrengthList": CreateLoraHookKeyframeFromStrengthList, + "ADE_SetLoraHookKeyframe": SetLoraHookKeyframes, + "ADE_AttachLoraHookToConditioning": SetModelLoraHook, "ADE_PairedConditioningSetMask": PairedConditioningSetMaskHooked, "ADE_ConditioningSetMask": ConditioningSetMaskHooked, "ADE_PairedConditioningSetMaskAndCombine": PairedConditioningSetMaskAndCombineHooked, @@ -168,8 +173,12 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_CombineLoraHooks": "Combine LoRA Hooks [2] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksFour": "Combine LoRA Hooks [4] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksEight": "Combine LoRA Hooks [8] πŸŽ­πŸ…πŸ…“", - "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", + "ADE_LoraHookKeyframe": "LoRA Hook Keyframe πŸŽ­πŸ…πŸ…“", + "ADE_LoraHookKeyframeInterpolation": "LoRA Hook Keyframes Interpolation πŸŽ­πŸ…πŸ…“", + "ADE_LoraHookKeyframeFromStrengthList": "LoRA Hook Keyframes From List πŸŽ­πŸ…πŸ…“", + "ADE_SetLoraHookKeyframe": "Set LoRA Hook Keyframes πŸŽ­πŸ…πŸ…“", + "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_PairedConditioningSetMask": "Set Props on Conds πŸŽ­πŸ…πŸ…“", "ADE_ConditioningSetMask": "Set Props on Cond πŸŽ­πŸ…πŸ…“", "ADE_PairedConditioningSetMaskAndCombine": "Set Props and Combine Conds πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 15b986f..59b18b4 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -2,6 +2,7 @@ import uuid import folder_paths from typing import Union from torch import Tensor +from collections.abc import Iterable from comfy.model_patcher import ModelPatcher from comfy.sd import CLIP @@ -9,8 +10,10 @@ import comfy.sd import comfy.utils from .conditioning import (COND_CONST, TimestepsCond, set_mask_conds, set_mask_and_combine_conds, set_unmasked_and_combine_conds, - LoraHook, LoraHookGroup) + LoraHook, LoraHookGroup, LoraHookKeyframe, LoraHookKeyframeGroup) from .model_injection import ModelPatcherAndInjector, CLIPWithHooks, load_hooked_lora_for_models, load_model_as_hooked_lora_for_models +from .utils_model import BIGMAX, InterpolationMethod +from .logger import logger ############################################### @@ -214,6 +217,150 @@ class ConditioningTimestepsNode: def create_schedule(self, start_percent: float, end_percent: float): return (TimestepsCond(start_percent=start_percent, end_percent=end_percent),) + +class SetLoraHookKeyframes: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "lora_hook": ("LORA_HOOK",), + "hook_kf": ("LORA_HOOK_KEYFRAMES",), + } + } + + RETURN_TYPES = ("LORA_HOOK",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "set_hook_keyframes" + + def set_hook_keyframes(self, lora_hook: LoraHookGroup, hook_kf: LoraHookKeyframeGroup): + new_lora_hook = lora_hook.clone() + new_lora_hook.set_keyframes_on_hooks(hook_kf=hook_kf) + return (new_lora_hook,) + + +class CreateLoraHookKeyframe: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}), + }, + "optional": { + "prev_hook_kf": ("LORA_HOOK_KEYFRAMES",), + } + } + + RETURN_TYPES = ("LORA_HOOK_KEYFRAMES",) + RETURN_NAMES = ("HOOK_KF",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/schedule lora hooks" + FUNCTION = "create_hook_keyframe" + + def create_hook_keyframe(self, strength_model: float, start_percent: float, guarantee_steps: float, + prev_hook_kf: LoraHookKeyframeGroup=None): + if prev_hook_kf: + prev_hook_kf = prev_hook_kf.clone() + else: + prev_hook_kf = LoraHookKeyframeGroup() + keyframe = LoraHookKeyframe(strength=strength_model, start_percent=start_percent, guarantee_steps=guarantee_steps) + prev_hook_kf.add(keyframe) + return (prev_hook_kf,) + + +class CreateLoraHookKeyframeInterpolation: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "strength_start": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "strength_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ), + "interpolation": (InterpolationMethod._LIST, ), + "intervals": ("INT", {"default": 5, "min": 2, "max": 100, "step": 1}), + "print_keyframes": ("BOOLEAN", {"default": False}), + }, + "optional": { + "prev_hook_kf": ("LORA_HOOK_KEYFRAMES",), + } + } + + RETURN_TYPES = ("LORA_HOOK_KEYFRAMES",) + RETURN_NAMES = ("HOOK_KF",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/schedule lora hooks" + FUNCTION = "create_hook_keyframes" + + def create_hook_keyframes(self, + start_percent: float, end_percent: float, + strength_start: float, strength_end: float, interpolation: str, intervals: int, + prev_hook_kf: LoraHookKeyframeGroup=None, print_keyframes=False): + if prev_hook_kf: + prev_hook_kf = prev_hook_kf.clone() + else: + prev_hook_kf = LoraHookKeyframeGroup() + percents = InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=intervals, method=interpolation) + strengths = InterpolationMethod.get_weights(num_from=strength_start, num_to=strength_end, length=intervals, method=interpolation) + + is_first = True + for percent, strength in zip(percents, strengths): + guarantee_steps = 0 + if is_first: + guarantee_steps = 1 + is_first = False + prev_hook_kf.add(LoraHookKeyframe(strength=strength, start_percent=percent, guarantee_steps=guarantee_steps)) + if print_keyframes: + logger.info(f"LoraHookKeyframe - start_percent:{percent} = {strength}") + return (prev_hook_kf,) + + +class CreateLoraHookKeyframeFromStrengthList: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "strengths_float": ("FLOAT", {"default": -1, "min": -1, "step": 0.001, "forceInput": True}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}), + "print_keyframes": ("BOOLEAN", {"default": False}), + }, + "optional": { + "prev_hook_kf": ("LORA_HOOK_KEYFRAMES",), + } + } + + RETURN_TYPES = ("LORA_HOOK_KEYFRAMES",) + RETURN_NAMES = ("HOOK_KF",) + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/schedule lora hooks" + FUNCTION = "create_hook_keyframes" + + def create_hook_keyframes(self, strengths_float: Union[float, list[float]], + start_percent: float, end_percent: float, + prev_hook_kf: LoraHookKeyframeGroup=None, print_keyframes=False): + if prev_hook_kf: + prev_hook_kf = prev_hook_kf.clone() + else: + prev_hook_kf = LoraHookKeyframeGroup() + if type(strengths_float) in (float, int): + strengths_float = [float(strengths_float)] + elif isinstance(strengths_float, Iterable): + pass + else: + raise Exception(f"strengths_floast must be either an interable input or a float, but was {type(strengths_float).__repr__}.") + percents = InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=len(strengths_float), method=InterpolationMethod.LINEAR) + + is_first = True + for percent, strength in zip(percents, strengths_float): + guarantee_steps = 0 + if is_first: + guarantee_steps = 1 + is_first = False + prev_hook_kf.add(LoraHookKeyframe(strength=strength, start_percent=percent, guarantee_steps=guarantee_steps)) + if print_keyframes: + logger.info(f"LoraHookKeyframe - start_percent:{percent} = {strength}") + return (prev_hook_kf,) + + ############################################### ############################################### ############################################### diff --git a/animatediff/sampling.py b/animatediff/sampling.py index 4af7771..c8d9351 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -41,8 +41,8 @@ class AnimateDiffHelper_GlobalState: self.params: InjectionParams = None self.sample_settings: SampleSettings = None self.reset() - - def initialize(self, model): + + def initialize(self, model: BaseModel): # this function is to be run in sampling func if not self.initialized: self.initialized = True @@ -53,15 +53,36 @@ class AnimateDiffHelper_GlobalState: if self.sample_settings.custom_cfg is not None: self.sample_settings.custom_cfg.initialize_timesteps(model) + def hooks_initialize(self, model: BaseModel, hook_groups: list[LoraHookGroup]): + # this function is to be run the first time all gathered + if not self.hooks_initialized: + self.hooks_initialized = True + for hook_group in hook_groups: + for hook in hook_group.hooks: + hook.reset() + hook.initialize_timesteps(model) + + def prepare_current_keyframes(self, timestep: Tensor): + if self.motion_models is not None: + self.motion_models.prepare_current_keyframe(t=timestep) + if self.params.context_options is not None: + self.params.context_options.prepare_current_context(t=timestep) + if self.sample_settings.custom_cfg is not None: + self.sample_settings.custom_cfg.prepare_current_keyframe(t=timestep) + + def prepare_hooks_current_keyframes(self, timestep: Tensor, hook_groups: list[LoraHookGroup]): + if self.model_patcher is not None: + self.model_patcher.prepare_hooked_patches_current_keyframe(t=timestep, hook_groups=hook_groups) + def reset(self): self.initialized = False + self.hooks_initialized = False self.start_step: int = 0 self.last_step: int = 0 self.current_step: int = 0 self.total_steps: int = 0 if self.model_patcher is not None: - self.model_patcher.unpatch_hooked() - self.model_patcher.clear_cached_hooked_weights() + self.model_patcher.clean_hooks() del self.model_patcher self.model_patcher = None if self.motion_models is not None: @@ -73,13 +94,13 @@ class AnimateDiffHelper_GlobalState: if self.sample_settings is not None: del self.sample_settings self.sample_settings = None - + def update_with_inject_params(self, params: InjectionParams): self.params = params def is_using_sliding_context(self): return self.params is not None and self.params.is_using_sliding_context() - + def create_exposed_params(self): # This dict will be exposed to be used by other extensions # DO NOT change any of the key names @@ -379,6 +400,7 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) -> ADGS.start_step = kwargs.get("start_step") or 0 ADGS.current_step = ADGS.start_step ADGS.last_step = kwargs.get("last_step") or 0 + ADGS.hooks_initialized = False if iter_opts.iterations > 1: logger.info(f"Iteration {curr_i+1}/{iter_opts.iterations}") # perform any iter_opts preprocessing on latents @@ -413,12 +435,7 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) -> def evolved_sampling_function(model, x, timestep, uncond, cond, cond_scale, model_options: dict={}, seed=None): ADGS.initialize(model) - if ADGS.motion_models is not None: - ADGS.motion_models.prepare_current_keyframe(t=timestep) - if ADGS.params.context_options is not None: - ADGS.params.context_options.prepare_current_context(t=timestep) - if ADGS.sample_settings.custom_cfg is not None: - ADGS.sample_settings.custom_cfg.prepare_current_keyframe(t=timestep) + ADGS.prepare_current_keyframes(timestep=timestep) # never use cfg1 optimization if using custom_cfg (since can have timesteps and such) if ADGS.sample_settings.custom_cfg is None and math.isclose(cond_scale, 1.0) and model_options.get("disable_cfg1_optimization", False) == False: @@ -513,11 +530,7 @@ def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options # look for control elif key == "control": control_item = cond_item - if hasattr(control_item, "sub_idxs"): - prepare_control_objects(control_item, full_idxs) - else: - raise ValueError(f"Control type {type(control_item).__name__} may not support required features for sliding context window; \ - use Control objects from Kosinkadink/ComfyUI-Advanced-ControlNet nodes, or make sure Advanced-ControlNet is updated.") + prepare_control_objects(control_item, full_idxs) resized_actual_cond[key] = control_item del control_item elif isinstance(cond_item, dict): @@ -616,17 +629,21 @@ def calc_cond_uncond_batch_wrapper(model, conds: list[dict], x_in: Tensor, times # check if conds or unconds contain lora_hook or default_cond contains_lora_hooks = False has_default_cond = False + hook_groups = [] for cond_uncond in conds: if cond_uncond is None: continue for t in cond_uncond: if COND_CONST.KEY_LORA_HOOK in t: contains_lora_hooks = True + hook_groups.append(t[COND_CONST.KEY_LORA_HOOK]) if COND_CONST.KEY_DEFAULT_COND in t: has_default_cond = True # if contains_lora_hooks: # break if contains_lora_hooks or has_default_cond: + ADGS.hooks_initialize(model, hook_groups=hook_groups) + ADGS.prepare_hooks_current_keyframes(timestep, hook_groups=hook_groups) return calc_conds_batch_lora_hook(model, conds, x_in, timestep, model_options, has_default_cond) # keep for backwards compatibility, for now if not hasattr(comfy.samplers, "calc_cond_batch"): From ba66cc1e1534ad7cf4993558d08d84178661342c Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Mon, 29 Apr 2024 11:13:05 -0500 Subject: [PATCH 28/30] Switched order of a couple nodes --- animatediff/nodes.py | 4 ++-- animatediff/nodes_conditioning.py | 2 -- 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/animatediff/nodes.py b/animatediff/nodes.py index ada1f0f..e9f9669 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -68,11 +68,11 @@ NODE_CLASS_MAPPINGS = { "ADE_CombineLoraHooks": CombineLoraHooks, "ADE_CombineLoraHooksFour": CombineLoraHookFourOptional, "ADE_CombineLoraHooksEight": CombineLoraHookEightOptional, + "ADE_SetLoraHookKeyframe": SetLoraHookKeyframes, "ADE_AttachLoraHookToCLIP": SetClipLoraHook, "ADE_LoraHookKeyframe": CreateLoraHookKeyframe, "ADE_LoraHookKeyframeInterpolation": CreateLoraHookKeyframeInterpolation, "ADE_LoraHookKeyframeFromStrengthList": CreateLoraHookKeyframeFromStrengthList, - "ADE_SetLoraHookKeyframe": SetLoraHookKeyframes, "ADE_AttachLoraHookToConditioning": SetModelLoraHook, "ADE_PairedConditioningSetMask": PairedConditioningSetMaskHooked, "ADE_ConditioningSetMask": ConditioningSetMaskHooked, @@ -173,11 +173,11 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ADE_CombineLoraHooks": "Combine LoRA Hooks [2] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksFour": "Combine LoRA Hooks [4] πŸŽ­πŸ…πŸ…“", "ADE_CombineLoraHooksEight": "Combine LoRA Hooks [8] πŸŽ­πŸ…πŸ…“", + "ADE_SetLoraHookKeyframe": "Set LoRA Hook Keyframes πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_LoraHookKeyframe": "LoRA Hook Keyframe πŸŽ­πŸ…πŸ…“", "ADE_LoraHookKeyframeInterpolation": "LoRA Hook Keyframes Interpolation πŸŽ­πŸ…πŸ…“", "ADE_LoraHookKeyframeFromStrengthList": "LoRA Hook Keyframes From List πŸŽ­πŸ…πŸ…“", - "ADE_SetLoraHookKeyframe": "Set LoRA Hook Keyframes πŸŽ­πŸ…πŸ…“", "ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook πŸŽ­πŸ…πŸ…“", "ADE_PairedConditioningSetMask": "Set Props on Conds πŸŽ­πŸ…πŸ…“", "ADE_ConditioningSetMask": "Set Props on Cond πŸŽ­πŸ…πŸ…“", diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 59b18b4..770e0ef 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -359,8 +359,6 @@ class CreateLoraHookKeyframeFromStrengthList: if print_keyframes: logger.info(f"LoraHookKeyframe - start_percent:{percent} = {strength}") return (prev_hook_kf,) - - ############################################### ############################################### ############################################### From fccdb43ea7263fda4af69cdf99e1bae64e3959f9 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 30 Apr 2024 02:18:03 -0500 Subject: [PATCH 29/30] Add sampler_post_cfg_function support for lora hooks --- animatediff/sampling.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/animatediff/sampling.py b/animatediff/sampling.py index c8d9351..90ba2c3 100644 --- a/animatediff/sampling.py +++ b/animatediff/sampling.py @@ -457,9 +457,8 @@ def evolved_sampling_function(model, x, timestep, uncond, cond, cond_scale, mode if hasattr(comfy.samplers, "cfg_function"): try: cached_calc_cond_batch = comfy.samplers.calc_cond_batch - # support sliding context for PAG/other sampler_post_cfg_function tech that may use calc_cond_batch - if ADGS.is_using_sliding_context(): - comfy.samplers.calc_cond_batch = wrapped_cfg_sliding_calc_cond_batch_factory(cached_calc_cond_batch) + # support hooks and sliding context for PAG/other sampler_post_cfg_function tech that may use calc_cond_batch + comfy.samplers.calc_cond_batch = wrapped_cfg_sliding_calc_cond_batch_factory(cached_calc_cond_batch) return comfy.samplers.cfg_function(model, cond_pred, uncond_pred, cond_scale, x, timestep, model_options, cond, uncond) finally: comfy.samplers.calc_cond_batch = cached_calc_cond_batch @@ -486,8 +485,10 @@ def wrapped_cfg_sliding_calc_cond_batch_factory(orig_calc_cond_batch): current_calc_cond_batch = comfy.samplers.calc_cond_batch # when inside sliding_calc_conds_batch, should return to original calc_cond_batch comfy.samplers.calc_cond_batch = orig_calc_cond_batch - result = sliding_calc_conds_batch(model, conds, x_in, timestep, model_options) - return result + if not ADGS.is_using_sliding_context(): + return calc_cond_uncond_batch_wrapper(model, conds, x_in, timestep, model_options) + else: + return sliding_calc_conds_batch(model, conds, x_in, timestep, model_options) finally: # make sure calc_cond_batch will become wrapped again comfy.samplers.calc_cond_batch = current_calc_cond_batch From 0a7ebbf7bd24093512cd054870f89beafd0ff21f Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Tue, 30 Apr 2024 08:55:26 -0500 Subject: [PATCH 30/30] Made all Combine LoRA Hook nodes act as a passthrough if only one hook provided instead of throwing an error --- animatediff/conditioning.py | 5 ++++- animatediff/nodes_conditioning.py | 7 +++++-- 2 files changed, 9 insertions(+), 3 deletions(-) diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py index 78dd23f..f626efb 100644 --- a/animatediff/conditioning.py +++ b/animatediff/conditioning.py @@ -91,13 +91,16 @@ class LoraHookGroup: hook.lora_keyframe = hook_kf @staticmethod - def combine_all_lora_hooks(lora_hooks_list: list['LoraHookGroup'], require_count=2) -> 'LoraHookGroup': + def combine_all_lora_hooks(lora_hooks_list: list['LoraHookGroup'], require_count=1) -> 'LoraHookGroup': actual: list[LoraHookGroup] = [] for group in lora_hooks_list: if group is not None: actual.append(group) if len(actual) < require_count: raise Exception(f"Need at least {require_count} LoRA Hooks to combine, but only had {len(actual)}.") + # if only 1 hook, just return itself without any cloning + if len(actual) == 1: + return actual[0] final_hook: LoraHookGroup = None for hook in actual: if final_hook is None: diff --git a/animatediff/nodes_conditioning.py b/animatediff/nodes_conditioning.py index 770e0ef..45bfc8c 100644 --- a/animatediff/nodes_conditioning.py +++ b/animatediff/nodes_conditioning.py @@ -545,6 +545,8 @@ class CombineLoraHooks: def INPUT_TYPES(s): return { "required": { + }, + "optional": { "lora_hook_A": ("LORA_HOOK",), "lora_hook_B": ("LORA_HOOK",), } @@ -554,8 +556,9 @@ class CombineLoraHooks: CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning/combine lora hooks" FUNCTION = "combine_lora_hooks" - def combine_lora_hooks(self, lora_hook_A: LoraHookGroup, lora_hook_B: LoraHookGroup): - return (lora_hook_A.clone_and_combine(lora_hook_B),) + def combine_lora_hooks(self, lora_hook_A: LoraHookGroup=None, lora_hook_B: LoraHookGroup=None): + candidates = [lora_hook_A, lora_hook_B] + return (LoraHookGroup.combine_all_lora_hooks(candidates),) class CombineLoraHookFourOptional: