Merge branch 'develop' into main-catchup

This commit is contained in:
Jedrzej Kosinski
2024-04-03 04:23:58 -05:00
committed by GitHub
11 changed files with 1057 additions and 84 deletions
View File
+555 -8
View File
@@ -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
+32 -13
View File
@@ -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) 🎭🅐🅓①",
}
+249
View File
@@ -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
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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
+1 -1
View File
@@ -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
-48
View File
@@ -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)
+1 -1
View File
@@ -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
+149 -9
View File
@@ -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
+66
View File
@@ -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):