diff --git a/animatediff/conditioning.py b/animatediff/conditioning.py new file mode 100644 index 0000000..e69de29 diff --git a/animatediff/model_injection.py b/animatediff/model_injection.py index a8c7886..affb13c 100644 --- a/animatediff/model_injection.py +++ b/animatediff/model_injection.py @@ -1,21 +1,25 @@ import copy -from typing import Union +from typing import Union, Callable from einops import rearrange from torch import Tensor 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 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, 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 @@ -41,12 +45,128 @@ class ModelPatcherAndInjector(ModelPatcher): 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.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 + # replace hook stuff + self.hooked_replace_patches = {} # binds LoraHook to specific replacement keys # injection stuff - self.motion_injection_params: InjectionParams = None + 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] + # 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 + 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 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 + + 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 model_patches_to(self, device): super().model_patches_to(device) if self.motion_models is not None: @@ -71,6 +191,9 @@ 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) + self.clear_cached_hooked_weights() return super().unpatch_model(device_to) else: return super().unpatch_model(device_to, unpatch_weights) @@ -78,6 +201,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) @@ -92,14 +216,436 @@ 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 + 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) + + 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() + + 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: + # 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) + 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() + + 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 + 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=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: + 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 + # 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 + 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 + # reinject model, if necessary + if was_injected: + self.inject_model() + + +class CLIPWithHooks(CLIP): + def __init__(self, clip: Union[CLIP, 'CLIPWithHooks']): + super().__init__(no_init=True) + 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.set_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 + 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) + + 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) + + +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 = 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 + 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 = {} + 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: 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) + else: + k = () + new_modelpatcher = None + + if clip is not None: + 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 + 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) + + +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 @@ -352,6 +898,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..4ef4bc2 100644 --- a/animatediff/nodes.py +++ b/animatediff/nodes.py @@ -6,6 +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, 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) @@ -17,7 +20,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 @@ -48,6 +51,16 @@ NODE_CLASS_MAPPINGS = { # Iteration Opts "ADE_IterationOptsDefault": IterationOptionsNode, "ADE_IterationOptsFreeInit": FreeInitOptionsNode, + # Conditioning + "ADE_RegisterLoraHook": MaskableLoraLoader, + "ADE_RegisterLoraHookModelOnly": MaskableLoraLoaderModelOnly, + #"ADE_RegisterModelAsLoraHook": MaskableSDModelLoader, # CLIP does not work properly on first run + "ADE_RegisterModelAsLoraHookModelOnly": MaskableSDModelLoaderModelOnly, + "ADE_CombineLoraHooks": CombineLoraHooks, + "ADE_CombineLoraHooksFour": CombineLoraHookFourOptional, + "ADE_CombineLoraHooksEight": CombineLoraHookEightOptional, + "ADE_AttachLoraHookToConditioning": SetModelLoraHook, + "ADE_AttachLoraHookToCLIP": SetClipLoraHook, # Noise Layer Nodes "ADE_NoiseLayerAdd": NoiseLayerAddNode, "ADE_NoiseLayerAddWeighted": NoiseLayerAddWeightedNode, @@ -78,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, @@ -92,12 +101,14 @@ 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, "ADE_AnimateDiffCombine": AnimateDiffCombine_Deprecated, + "ADE_AnimateDiffModelSettings_Release": AnimateDiffModelSettings, + "ADE_AnimateDiffModelSettingsSimple": AnimateDiffModelSettingsSimple, + "ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced, + "ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": AnimateDiffModelSettingsAdvancedAttnStrengths, } NODE_DISPLAY_NAME_MAPPINGS = { # Unencapsulated @@ -121,6 +132,16 @@ NODE_DISPLAY_NAME_MAPPINGS = { # Iteration Opts "ADE_IterationOptsDefault": "Default Iteration Options πŸŽ­πŸ…πŸ…“", "ADE_IterationOptsFreeInit": "FreeInit Iteration Options πŸŽ­πŸ…πŸ…“", + # 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_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 πŸŽ­πŸ…πŸ…“", # Noise Layer Nodes "ADE_NoiseLayerAdd": "Noise Layer [Add] πŸŽ­πŸ…πŸ…“", "ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] πŸŽ­πŸ…πŸ…“", @@ -151,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 πŸŽ­πŸ…πŸ…“β‘‘", @@ -165,10 +182,12 @@ 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] πŸŽ­πŸ…πŸ…“", "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 new file mode 100644 index 0000000..d759fb7 --- /dev/null +++ b/animatediff/nodes_conditioning.py @@ -0,0 +1,249 @@ +import uuid +import folder_paths +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, load_model_as_hooked_lora_for_models + +# 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_HOOK") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "load_lora" + + 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) + + 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) + + 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): + @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_HOOK") + CATEGORY = "Animate Diff πŸŽ­πŸ…πŸ…“/conditioning" + FUNCTION = "load_lora_model_only" + + def load_lora_model_only(self, model: ModelPatcher, lora_name: str, strength_model: float): + 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 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): + 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 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: + @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/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 { + "required": { + }, + "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/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/nodes_lora.py b/animatediff/nodes_lora.py index 5cc3ba3..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 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 47888e6..b88e432 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,11 @@ 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: + 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: del self.motion_models self.motion_models = None @@ -306,6 +312,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,10 +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(): - if hasattr(comfy.samplers, "calc_cond_batch"): - cond_pred, uncond_pred = comfy.samplers.calc_cond_batch(model, [cond, uncond_], x, timestep, model_options) - else: - 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) @@ -424,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): @@ -519,10 +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 - if hasattr(comfy.samplers, "calc_cond_batch"): - sub_cond_out, sub_uncond_out = comfy.samplers.calc_cond_batch(model, [sub_cond, sub_uncond], sub_x, sub_timestep, model_options) - else: - 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 @@ -557,3 +559,141 @@ 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, 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 conds: + if cond_uncond is None: + continue + 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, 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) + 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..5b8f1f6 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,71 @@ class ADKeyframeGroup: return cloned +class LoraHookMode: + MIN_VRAM = "min_vram" + MIN_VRAM_LOWVRAM = "min_vram_lowvram" + MAX_SPEED = "max_speed" + 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):