Merge PR #330 from Kosinkadink/develop - SD LoRA masking and scheduling
Add SD LoRA masking, scheduling, and conditioning helpers.
This commit is contained in:
@@ -0,0 +1,303 @@
|
||||
from torch import Tensor
|
||||
|
||||
from comfy.model_base import BaseModel
|
||||
|
||||
from .utils_motion import get_sorted_list_via_attr
|
||||
|
||||
|
||||
class LoraHookMode:
|
||||
MIN_VRAM = "min_vram"
|
||||
MAX_SPEED = "max_speed"
|
||||
#MIN_VRAM_LOWVRAM = "min_vram_lowvram"
|
||||
#MAX_SPEED_LOWVRAM = "max_speed_lowvram"
|
||||
|
||||
|
||||
# Acts simply as a way to track unique LoraHooks
|
||||
class HookRef:
|
||||
pass
|
||||
|
||||
|
||||
class LoraHook:
|
||||
def __init__(self, lora_name: str):
|
||||
self.lora_name = lora_name
|
||||
self.lora_keyframe = LoraHookKeyframeGroup()
|
||||
self.hook_ref = HookRef()
|
||||
|
||||
def initialize_timesteps(self, model: BaseModel):
|
||||
self.lora_keyframe.initialize_timesteps(model)
|
||||
|
||||
def reset(self):
|
||||
self.lora_keyframe.reset()
|
||||
|
||||
|
||||
def get_copy(self):
|
||||
'''
|
||||
Copies LoraHook, but maintains same HookRef
|
||||
'''
|
||||
c = LoraHook(lora_name=self.lora_name)
|
||||
c.lora_keyframe = self.lora_keyframe
|
||||
c.hook_ref = self.hook_ref # same instance that acts as ref
|
||||
return c
|
||||
|
||||
@property
|
||||
def strength(self):
|
||||
return self.lora_keyframe.strength
|
||||
|
||||
def __eq__(self, other: 'LoraHook'):
|
||||
return self.__class__ == other.__class__ and self.hook_ref == other.hook_ref
|
||||
|
||||
def __hash__(self):
|
||||
return hash(self.hook_ref)
|
||||
|
||||
|
||||
class LoraHookGroup:
|
||||
'''
|
||||
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: LoraHook):
|
||||
if hook not in self.hooks:
|
||||
self.hooks.append(hook)
|
||||
|
||||
def is_empty(self):
|
||||
return len(self.hooks) == 0
|
||||
|
||||
def contains(self, lora_hook: LoraHook):
|
||||
return lora_hook in self.hooks
|
||||
|
||||
def clone(self):
|
||||
cloned = LoraHookGroup()
|
||||
for hook in self.hooks:
|
||||
cloned.add(hook.get_copy())
|
||||
return cloned
|
||||
|
||||
def clone_and_combine(self, other: 'LoraHookGroup'):
|
||||
cloned = self.clone()
|
||||
for hook in other.hooks:
|
||||
cloned.add(hook.get_copy())
|
||||
return cloned
|
||||
|
||||
def set_keyframes_on_hooks(self, hook_kf: 'LoraHookKeyframeGroup'):
|
||||
hook_kf = hook_kf.clone()
|
||||
for hook in self.hooks:
|
||||
hook.lora_keyframe = hook_kf
|
||||
|
||||
@staticmethod
|
||||
def combine_all_lora_hooks(lora_hooks_list: list['LoraHookGroup'], require_count=1) -> 'LoraHookGroup':
|
||||
actual: list[LoraHookGroup] = []
|
||||
for group in lora_hooks_list:
|
||||
if group is not None:
|
||||
actual.append(group)
|
||||
if len(actual) < require_count:
|
||||
raise Exception(f"Need at least {require_count} LoRA Hooks to combine, but only had {len(actual)}.")
|
||||
# if only 1 hook, just return itself without any cloning
|
||||
if len(actual) == 1:
|
||||
return actual[0]
|
||||
final_hook: LoraHookGroup = None
|
||||
for hook in actual:
|
||||
if final_hook is None:
|
||||
final_hook = hook.clone()
|
||||
else:
|
||||
final_hook = final_hook.clone_and_combine(hook)
|
||||
return final_hook
|
||||
|
||||
|
||||
class LoraHookKeyframe:
|
||||
def __init__(self, strength: float, start_percent=0.0, guarantee_steps=1):
|
||||
self.strength = strength
|
||||
# scheduling
|
||||
self.start_percent = float(start_percent)
|
||||
self.start_t = 999999999.9
|
||||
self.guarantee_steps = guarantee_steps
|
||||
|
||||
def clone(self):
|
||||
c = LoraHookKeyframe(strength=self.strength,
|
||||
start_percent=self.start_percent, guarantee_steps=self.guarantee_steps)
|
||||
c.start_t = self.start_t
|
||||
return c
|
||||
|
||||
class LoraHookKeyframeGroup:
|
||||
def __init__(self):
|
||||
self.keyframes: list[LoraHookKeyframe] = []
|
||||
self._current_keyframe: LoraHookKeyframe = None
|
||||
self._current_used_steps: int = 0
|
||||
self._current_index: int = 0
|
||||
self._curr_t: float = -1
|
||||
|
||||
def reset(self):
|
||||
self._current_keyframe = None
|
||||
self._current_used_steps = 0
|
||||
self._current_index = 0
|
||||
self._curr_t = -1
|
||||
self._set_first_as_current()
|
||||
|
||||
def add(self, keyframe: LoraHookKeyframe):
|
||||
# add to end of list, then sort
|
||||
self.keyframes.append(keyframe)
|
||||
self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent")
|
||||
self._set_first_as_current()
|
||||
|
||||
def _set_first_as_current(self):
|
||||
if len(self.keyframes) > 0:
|
||||
self._current_keyframe = self.keyframes[0]
|
||||
else:
|
||||
self._current_keyframe = None
|
||||
|
||||
def has_index(self, index: int) -> int:
|
||||
return index >= 0 and index < len(self.keyframes)
|
||||
|
||||
def is_empty(self) -> bool:
|
||||
return len(self.keyframes) == 0
|
||||
|
||||
def clone(self):
|
||||
cloned = LoraHookKeyframeGroup()
|
||||
for keyframe in self.keyframes:
|
||||
cloned.keyframes.append(keyframe)
|
||||
cloned._set_first_as_current()
|
||||
return cloned
|
||||
|
||||
def initialize_timesteps(self, model: BaseModel):
|
||||
for keyframe in self.keyframes:
|
||||
keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent)
|
||||
|
||||
def prepare_current_keyframe(self, curr_t: float) -> bool:
|
||||
if self.is_empty():
|
||||
return False
|
||||
if curr_t == self._curr_t:
|
||||
return False
|
||||
prev_index = self._current_index
|
||||
# if met guaranteed steps, look for next keyframe in case need to switch
|
||||
if self._current_used_steps >= self._current_keyframe.guarantee_steps:
|
||||
# if has next index, loop through and see if need t oswitch
|
||||
if self.has_index(self._current_index+1):
|
||||
for i in range(self._current_index+1, len(self.keyframes)):
|
||||
eval_c = self.keyframes[i]
|
||||
# check if start_t is greater or equal to curr_t
|
||||
# NOTE: t is in terms of sigmas, not percent, so bigger number = earlier step in sampling
|
||||
if eval_c.start_t >= curr_t:
|
||||
self._current_index = i
|
||||
self._current_keyframe = eval_c
|
||||
self._current_used_steps = 0
|
||||
# if guarantee_steps greater than zero, stop searching for other keyframes
|
||||
if self._current_keyframe.guarantee_steps > 0:
|
||||
break
|
||||
# if eval_c is outside the percent range, stop looking further
|
||||
else: break
|
||||
# update steps current context is used
|
||||
self._current_used_steps += 1
|
||||
# update current timestep this was performed on
|
||||
self._curr_t = curr_t
|
||||
# return True if keyframe changed, False if no change
|
||||
return prev_index != self._current_index
|
||||
|
||||
# properties shadow those of LoraHookKeyframe
|
||||
@property
|
||||
def strength(self):
|
||||
if self._current_keyframe is not None:
|
||||
return self._current_keyframe.strength
|
||||
return 1.0
|
||||
|
||||
|
||||
class COND_CONST:
|
||||
KEY_LORA_HOOK = "lora_hook"
|
||||
KEY_DEFAULT_COND = "default_cond"
|
||||
|
||||
COND_AREA_DEFAULT = "default"
|
||||
COND_AREA_MASK_BOUNDS = "mask bounds"
|
||||
_LIST_COND_AREA = [COND_AREA_DEFAULT, COND_AREA_MASK_BOUNDS]
|
||||
|
||||
|
||||
class TimestepsCond:
|
||||
def __init__(self, start_percent: float, end_percent: float):
|
||||
self.start_percent = start_percent
|
||||
self.end_percent = end_percent
|
||||
|
||||
|
||||
def conditioning_set_values(conditioning, values={}):
|
||||
c = []
|
||||
for t in conditioning:
|
||||
n = [t[0], t[1].copy()]
|
||||
for k in values:
|
||||
n[1][k] = values[k]
|
||||
c.append(n)
|
||||
return c
|
||||
|
||||
def set_lora_hook_for_conditioning(conditioning, lora_hook: LoraHookGroup):
|
||||
if lora_hook is None:
|
||||
return conditioning
|
||||
return conditioning_set_values(conditioning, {COND_CONST.KEY_LORA_HOOK: lora_hook})
|
||||
|
||||
def set_timesteps_for_conditioning(conditioning, timesteps_cond: TimestepsCond):
|
||||
if timesteps_cond is None:
|
||||
return conditioning
|
||||
return conditioning_set_values(conditioning, {"start_percent": timesteps_cond.start_percent,
|
||||
"end_percent": timesteps_cond.end_percent})
|
||||
|
||||
def set_mask_for_conditioning(conditioning, mask: Tensor, set_cond_area: str, strength: float):
|
||||
if mask is None:
|
||||
return conditioning
|
||||
set_area_to_bounds = False
|
||||
if set_cond_area != COND_CONST.COND_AREA_DEFAULT:
|
||||
set_area_to_bounds = True
|
||||
if len(mask.shape) < 3:
|
||||
mask = mask.unsqueeze(0)
|
||||
|
||||
return conditioning_set_values(conditioning, {"mask": mask,
|
||||
"set_area_to_bounds": set_area_to_bounds,
|
||||
"mask_strength": strength})
|
||||
|
||||
def combine_conditioning(conds: list):
|
||||
combined_conds = []
|
||||
for cond in conds:
|
||||
combined_conds.extend(cond)
|
||||
return combined_conds
|
||||
|
||||
def set_mask_conds(conds: list, strength: float, set_cond_area: str,
|
||||
opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None):
|
||||
masked_conds = []
|
||||
for c in conds:
|
||||
# first, apply lora_hook to conditioning, if provided
|
||||
c = set_lora_hook_for_conditioning(c, opt_lora_hook)
|
||||
# next, apply mask to conditioning
|
||||
c = set_mask_for_conditioning(conditioning=c, mask=opt_mask, strength=strength, set_cond_area=set_cond_area)
|
||||
# apply timesteps, if present
|
||||
c = set_timesteps_for_conditioning(conditioning=c, timesteps_cond=opt_timesteps)
|
||||
# finally, apply mask to conditioning and store
|
||||
masked_conds.append(c)
|
||||
return masked_conds
|
||||
|
||||
def set_mask_and_combine_conds(conds: list, new_conds: list, strength: float, set_cond_area: str,
|
||||
opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None):
|
||||
combined_conds = []
|
||||
for c, masked_c in zip(conds, new_conds):
|
||||
# first, apply lora_hook to new conditioning, if provided
|
||||
masked_c = set_lora_hook_for_conditioning(masked_c, opt_lora_hook)
|
||||
# next, apply mask to new conditioning, if provided
|
||||
masked_c = set_mask_for_conditioning(conditioning=masked_c, mask=opt_mask, set_cond_area=set_cond_area, strength=strength)
|
||||
# apply timesteps, if present
|
||||
masked_c = set_timesteps_for_conditioning(conditioning=masked_c, timesteps_cond=opt_timesteps)
|
||||
# finally, combine with existing conditioning and store
|
||||
combined_conds.append(combine_conditioning([c, masked_c]))
|
||||
return combined_conds
|
||||
|
||||
def set_unmasked_and_combine_conds(conds: list, new_conds: list,
|
||||
opt_lora_hook: LoraHookGroup, opt_timesteps: TimestepsCond=None):
|
||||
combined_conds = []
|
||||
for c, new_c in zip(conds, new_conds):
|
||||
# first, apply lora_hook to new conditioning, if provided
|
||||
new_c = set_lora_hook_for_conditioning(new_c, opt_lora_hook)
|
||||
# next, add default_cond key to cond so that during sampling, it can be identified
|
||||
new_c = conditioning_set_values(new_c, {COND_CONST.KEY_DEFAULT_COND: True})
|
||||
# apply timesteps, if present
|
||||
new_c = set_timesteps_for_conditioning(conditioning=new_c, timesteps_cond=opt_timesteps)
|
||||
# finally, combine with existing conditioning and store
|
||||
combined_conds.append(combine_conditioning([c, new_c]))
|
||||
return combined_conds
|
||||
+582
-13
@@ -1,15 +1,19 @@
|
||||
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 math
|
||||
|
||||
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 .adapter_cameractrl import CameraPoseEncoder, CameraEntry, prepare_pose_embedding
|
||||
@@ -18,6 +22,7 @@ from .motion_module_ad import (AnimateDiffModel, AnimateDiffFormat, EncoderOnlyA
|
||||
has_mid_block, normalize_ad_state_dict, get_position_encoding_max_len)
|
||||
from .logger import logger
|
||||
from .utils_motion import ADKeyframe, ADKeyframeGroup, MotionCompatibilityError, get_combined_multival, ade_broadcast_image_to, normalize_min_max
|
||||
from .conditioning import HookRef, LoraHook, LoraHookGroup, LoraHookMode
|
||||
from .motion_lora import MotionLoraInfo, MotionLoraList
|
||||
from .utils_model import get_motion_lora_path, get_motion_model_path, get_sd_model_type
|
||||
from .sample_settings import SampleSettings, SeedNoiseGeneration
|
||||
@@ -43,12 +48,146 @@ class ModelPatcherAndInjector(ModelPatcher):
|
||||
if hasattr(m, "object_patches_backup"):
|
||||
self.object_patches_backup = m.object_patches_backup
|
||||
|
||||
|
||||
# lora hook stuff
|
||||
self.hooked_patches: dict[HookRef] = {} # binds LoraHook to specific keys
|
||||
self.hooked_backup: dict[str, tuple[Tensor, torch.device]] = {}
|
||||
self.cached_hooked_patches: dict[LoraHookGroup, dict[str, Tensor]] = {} # binds LoraHookGroup to pre-calculated weights (speed optimization)
|
||||
self.current_lora_hooks = None
|
||||
self.lora_hook_mode = LoraHookMode.MAX_SPEED
|
||||
self.model_params_lowvram = False
|
||||
self.model_params_lowvram_keys = {} # keeps track of keys with applied 'weight_function' or 'bias_function'
|
||||
# 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_ref in self.hooked_patches:
|
||||
cloned.hooked_patches[hook_ref] = {}
|
||||
for k in self.hooked_patches[hook_ref]:
|
||||
cloned.hooked_patches[hook_ref][k] = self.hooked_patches[hook_ref][k][:]
|
||||
# copy pre-calc weights bound to LoraHookGroups
|
||||
for group in self.cached_hooked_patches:
|
||||
cloned.cached_hooked_patches[group] = {}
|
||||
for k in self.cached_hooked_patches[group]:
|
||||
cloned.cached_hooked_patches[group][k] = self.cached_hooked_patches[group][k]
|
||||
cloned.hooked_backup = self.hooked_backup
|
||||
cloned.current_lora_hooks = self.current_lora_hooks
|
||||
cloned.currently_injected = self.currently_injected
|
||||
cloned.lora_hook_mode = self.lora_hook_mode
|
||||
if not hooks_only:
|
||||
cloned.motion_models = self.motion_models.clone() if self.motion_models else self.motion_models
|
||||
cloned.sample_settings = self.sample_settings
|
||||
cloned.motion_injection_params = self.motion_injection_params.clone() if self.motion_injection_params else self.motion_injection_params
|
||||
return cloned
|
||||
|
||||
@classmethod
|
||||
def create_from(cls, model: Union[ModelPatcher, 'ModelPatcherAndInjector'], hooks_only=False) -> 'ModelPatcherAndInjector':
|
||||
if isinstance(model, ModelPatcherAndInjector):
|
||||
return model.clone(hooks_only=hooks_only)
|
||||
else:
|
||||
return ModelPatcherAndInjector(model)
|
||||
|
||||
def 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:
|
||||
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
|
||||
return returned
|
||||
|
||||
def set_lora_hook_mode(self, lora_hook_mode: str):
|
||||
self.lora_hook_mode = lora_hook_mode
|
||||
|
||||
def prepare_hooked_patches_current_keyframe(self, t: Tensor, hook_groups: list[LoraHookGroup]):
|
||||
curr_t = t[0]
|
||||
for hook_group in hook_groups:
|
||||
for hook in hook_group.hooks:
|
||||
changed = hook.lora_keyframe.prepare_current_keyframe(curr_t=curr_t)
|
||||
# if keyframe changed, remove any cached LoraHookGroups that contain hook with the same hook_ref;
|
||||
# this will cause the weights to be recalculated when sampling
|
||||
if changed:
|
||||
for cached_group in list(self.cached_hooked_patches.keys()):
|
||||
if cached_group.contains(hook):
|
||||
self.cached_hooked_patches.pop(cached_group)
|
||||
|
||||
def clean_hooks(self):
|
||||
self.unpatch_hooked()
|
||||
self.clear_cached_hooked_weights()
|
||||
# for lora_hook in self.hooked_patches:
|
||||
# lora_hook.reset()
|
||||
|
||||
def add_hooked_patches(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0):
|
||||
'''
|
||||
Based on add_patches, but for hooked weights.
|
||||
'''
|
||||
# TODO: make this work with timestep scheduling
|
||||
current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook.hook_ref, {})
|
||||
p = set()
|
||||
for key in patches:
|
||||
if key in self.model_keys:
|
||||
p.add(key)
|
||||
current_patches: list[tuple] = current_hooked_patches.get(key, [])
|
||||
current_patches.append((strength_patch, patches[key], strength_model))
|
||||
current_hooked_patches[key] = current_patches
|
||||
self.hooked_patches[lora_hook.hook_ref] = current_hooked_patches
|
||||
# since should care about these patches too to determine if same model, reroll patches_uuid
|
||||
self.patches_uuid = uuid.uuid4()
|
||||
return list(p)
|
||||
|
||||
def add_hooked_patches_as_diffs(self, lora_hook: LoraHook, patches: dict, strength_patch=1.0, strength_model=1.0):
|
||||
'''
|
||||
Based on add_hooked_patches, but intended for using a model's weights as lora hook.
|
||||
'''
|
||||
# TODO: make this work with timestep scheduling
|
||||
current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook.hook_ref, {})
|
||||
p = set()
|
||||
for key in patches:
|
||||
if key in self.model_keys:
|
||||
p.add(key)
|
||||
current_patches: list[tuple] = current_hooked_patches.get(key, [])
|
||||
# take difference between desired weight and existing weight to get diff
|
||||
current_patches.append((strength_patch, (patches[key]-comfy.utils.get_attr(self.model, key),), strength_model))
|
||||
current_hooked_patches[key] = current_patches
|
||||
self.hooked_patches[lora_hook.hook_ref] = current_hooked_patches
|
||||
# since should care about these patches too to determine if same model, reroll patches_uuid
|
||||
self.patches_uuid = uuid.uuid4()
|
||||
return list(p)
|
||||
|
||||
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.hook_ref, {})
|
||||
for key in hook_patches.keys():
|
||||
current_patches: list[tuple] = combined_patches.get(key, [])
|
||||
if math.isclose(hook.strength, 1.0):
|
||||
# if hook strength is 1.0, can just add it directly
|
||||
current_patches.extend(hook_patches[key])
|
||||
else:
|
||||
# otherwise, need to multiply original patch strength by hook strength
|
||||
# patches are stored as tuples: (strength_patch, (tuple_with_weights,), strength_model)
|
||||
for patch in hook_patches[key]:
|
||||
new_patch = list(patch)
|
||||
new_patch[0] *= hook.strength
|
||||
current_patches.append(tuple(new_patch))
|
||||
combined_patches[key] = current_patches
|
||||
return combined_patches
|
||||
|
||||
def model_patches_to(self, device):
|
||||
super().model_patches_to(device)
|
||||
|
||||
@@ -59,35 +198,464 @@ class ModelPatcherAndInjector(ModelPatcher):
|
||||
else:
|
||||
patched_model = super().patch_model(device_to, patch_weights)
|
||||
# finally, perform motion model injection
|
||||
self.inject_model(device_to=device_to)
|
||||
self.inject_model()
|
||||
return patched_model
|
||||
|
||||
def patch_model_lowvram(self, *args, **kwargs):
|
||||
try:
|
||||
return super().patch_model_lowvram(*args, **kwargs)
|
||||
finally:
|
||||
# check if any modules have weight_function or bias_function that is not None
|
||||
# NOTE: this serves no purpose currently, but I have it here for future reasons
|
||||
for n, m in self.model.named_modules():
|
||||
if not hasattr(m, "comfy_cast_weights"):
|
||||
continue
|
||||
if getattr(m, "weight_function", None) is not None:
|
||||
self.model_params_lowvram = True
|
||||
self.model_params_lowvram_keys[f"{n}.weight"] = n
|
||||
if getattr(m, "bias_function", None) is not None:
|
||||
self.model_params_lowvram = True
|
||||
self.model_params_lowvram_keys[f"{n}.weight"] = n
|
||||
|
||||
def unpatch_model(self, device_to=None, unpatch_weights=True):
|
||||
# first, eject motion model from unet
|
||||
self.eject_model(device_to=device_to)
|
||||
self.eject_model()
|
||||
# finally, do normal model unpatching
|
||||
if unpatch_weights: # TODO: keep only 'else' portion when don't need to worry about past comfy versions
|
||||
return super().unpatch_model(device_to)
|
||||
# handle hooked_patches first
|
||||
self.clean_hooks()
|
||||
try:
|
||||
return super().unpatch_model(device_to)
|
||||
finally:
|
||||
self.model_params_lowvram = False
|
||||
self.model_params_lowvram_keys.clear()
|
||||
else:
|
||||
return super().unpatch_model(device_to, unpatch_weights)
|
||||
try:
|
||||
return super().unpatch_model(device_to, unpatch_weights)
|
||||
finally:
|
||||
self.model_params_lowvram = False
|
||||
self.model_params_lowvram_keys.clear()
|
||||
|
||||
def inject_model(self, device_to=None):
|
||||
def inject_model(self):
|
||||
if self.motion_models is not None:
|
||||
for motion_model in self.motion_models.models:
|
||||
self.currently_injected = True
|
||||
motion_model.model.inject(self)
|
||||
|
||||
def eject_model(self, device_to=None):
|
||||
def eject_model(self):
|
||||
if self.motion_models is not None:
|
||||
for motion_model in self.motion_models.models:
|
||||
motion_model.model.eject(self)
|
||||
self.currently_injected = False
|
||||
|
||||
def apply_lora_hooks(self, lora_hooks: LoraHookGroup):
|
||||
# first, determine if need to reapply patches
|
||||
if self.current_lora_hooks == lora_hooks:
|
||||
return
|
||||
# patch hooks
|
||||
self.patch_hooked(lora_hooks=lora_hooks)
|
||||
|
||||
def patch_hooked(self, lora_hooks: LoraHookGroup) -> None:
|
||||
# first, unpatch any previous patches
|
||||
self.unpatch_hooked()
|
||||
# eject model, if needed
|
||||
was_injected = self.currently_injected
|
||||
if was_injected:
|
||||
self.eject_model()
|
||||
|
||||
model_sd = self.model_state_dict()
|
||||
# if have cached weights for lora_hooks, use it
|
||||
cached_weights = self.cached_hooked_patches.get(lora_hooks, None)
|
||||
if cached_weights is not None:
|
||||
for key in cached_weights:
|
||||
if key not in model_sd:
|
||||
logger.warning(f"Cached LoraHook could not patch. key doesn't exist in model: {key}")
|
||||
self.patch_cached_hooked_weight(cached_weights=cached_weights, key=key)
|
||||
else:
|
||||
# get combined patches of relevant lora_hooks
|
||||
relevant_patches = self.get_combined_hooked_patches(lora_hooks=lora_hooks)
|
||||
for key in relevant_patches:
|
||||
if key not in model_sd:
|
||||
logger.warning(f"LoraHook could not patch. key doesn't exist in model: {key}")
|
||||
continue
|
||||
self.patch_hooked_weight_to_device(lora_hooks=lora_hooks, combined_patches=relevant_patches, key=key)
|
||||
self.current_lora_hooks = lora_hooks
|
||||
# reinject model, if needed
|
||||
if was_injected:
|
||||
self.inject_model()
|
||||
|
||||
def patch_cached_hooked_weight(self, cached_weights: dict, key: str):
|
||||
# TODO: handle model_params_lowvram stuff if necessary
|
||||
inplace_update = self.weight_inplace_update
|
||||
if key not in self.hooked_backup:
|
||||
weight: Tensor = comfy.utils.get_attr(self.model, key)
|
||||
target_device = self.offload_device
|
||||
if self.lora_hook_mode == LoraHookMode.MAX_SPEED:
|
||||
target_device = weight.device
|
||||
self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device)
|
||||
if inplace_update:
|
||||
comfy.utils.copy_to_param(self.model, key, cached_weights[key])
|
||||
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):
|
||||
if key not in combined_patches:
|
||||
return
|
||||
|
||||
inplace_update = self.weight_inplace_update
|
||||
weight: Tensor = comfy.utils.get_attr(self.model, key)
|
||||
if key not in self.hooked_backup:
|
||||
target_device = self.offload_device
|
||||
if self.lora_hook_mode == LoraHookMode.MAX_SPEED:
|
||||
target_device = weight.device
|
||||
self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device)
|
||||
|
||||
# TODO: handle model_params_lowvram stuff if necessary
|
||||
temp_weight = comfy.model_management.cast_to_device(weight, weight.device, 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):
|
||||
# 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
|
||||
|
||||
inplace_update = self.weight_inplace_update
|
||||
weight: Tensor = comfy.utils.get_attr(self.model, key)
|
||||
if key not in self.hooked_backup:
|
||||
# TODO: handle model_params_lowvram stuff if necessary
|
||||
target_device = self.offload_device
|
||||
if self.lora_hook_mode == LoraHookMode.MAX_SPEED:
|
||||
target_device = weight.device
|
||||
self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device)
|
||||
|
||||
out_weight = replace_patches[key].to(weight.device)
|
||||
if self.lora_hook_mode == LoraHookMode.MAX_SPEED:
|
||||
self.cached_hooked_patches.setdefault(lora_hooks, {})
|
||||
self.cached_hooked_patches[lora_hooks][key] = out_weight
|
||||
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) -> None:
|
||||
# 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 model_params_lowvram stuff if necessary
|
||||
keys = list(self.hooked_backup.keys())
|
||||
if self.weight_inplace_update:
|
||||
for k in keys:
|
||||
if self.lora_hook_mode == LoraHookMode.MAX_SPEED: # does not need to be casted - cache device matches needed device
|
||||
comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k][0])
|
||||
else: # should be casted as may not match needed device
|
||||
comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1]))
|
||||
else:
|
||||
for k in keys:
|
||||
if self.lora_hook_mode == LoraHookMode.MAX_SPEED:
|
||||
comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k][0])
|
||||
else: # should be casted as may not match needed device
|
||||
comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1]))
|
||||
# clear hooked_backup
|
||||
self.hooked_backup.clear()
|
||||
self.current_lora_hooks = None
|
||||
# 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_patches_as_diffs(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0):
|
||||
return self.patcher.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=patches, strength_patch=strength_patch, strength_model=strength_model)
|
||||
|
||||
|
||||
class ModelPatcherCLIPHooks(ModelPatcher):
|
||||
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: dict[str, tuple[Tensor, torch.device]] = {}
|
||||
|
||||
self.current_lora_hooks = None
|
||||
self.desired_lora_hooks = None
|
||||
self.lora_hook_mode = LoraHookMode.MAX_SPEED
|
||||
|
||||
self.model_params_lowvram = False
|
||||
self.model_params_lowvram_keys = {} # keeps track of keys with applied 'weight_function' or 'bias_function'
|
||||
|
||||
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][:]
|
||||
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
|
||||
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.
|
||||
'''
|
||||
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 add_hooked_patches_as_diffs(self, lora_hook: LoraHook, patches, strength_patch=1.0, strength_model=1.0):
|
||||
'''
|
||||
Based on add_hooked_patches, but intended for using a model's weights as lora hook.
|
||||
'''
|
||||
current_hooked_patches: dict[str,list] = self.hooked_patches.get(lora_hook, {})
|
||||
p = set()
|
||||
for key in patches:
|
||||
if key in self.model_keys:
|
||||
p.add(key)
|
||||
current_patches: list[tuple] = current_hooked_patches.get(key, [])
|
||||
# take difference between desired weight and existing weight to get diff
|
||||
current_patches.append((strength_patch, (patches[key]-comfy.utils.get_attr(self.model, key),), strength_model))
|
||||
current_hooked_patches[key] = current_patches
|
||||
self.hooked_patches[lora_hook] = current_hooked_patches
|
||||
# since should care about these patches too to determine if same model, reroll patches_uuid
|
||||
self.patches_uuid = uuid.uuid4()
|
||||
return list(p)
|
||||
|
||||
def get_combined_hooked_patches(self, lora_hooks: LoraHookGroup):
|
||||
'''
|
||||
Returns patches for selected lora_hooks.
|
||||
'''
|
||||
# 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 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 = weight.device
|
||||
|
||||
if key not in self.hooked_backup:
|
||||
self.hooked_backup[key] = (weight.to(device=target_device, copy=inplace_update), weight.device)
|
||||
out_weight = replace_patches[key].to(target_device)
|
||||
if inplace_update:
|
||||
comfy.utils.copy_to_param(self.model, key, out_weight)
|
||||
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()
|
||||
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 patch_model_lowvram(self, *args, **kwargs):
|
||||
try:
|
||||
return super().patch_model_lowvram(*args, **kwargs)
|
||||
finally:
|
||||
# check if any modules have weight_function or bias_function that is not None
|
||||
# NOTE: this serves no purpose currently, but I have it here for future reasons
|
||||
for n, m in self.model.named_modules():
|
||||
if not hasattr(m, "comfy_cast_weights"):
|
||||
continue
|
||||
if getattr(m, "weight_function", None) is not None:
|
||||
self.model_params_lowvram = True
|
||||
self.model_params_lowvram_keys[f"{n}.weight"] = n
|
||||
if getattr(m, "bias_function", None) is not None:
|
||||
self.model_params_lowvram = True
|
||||
self.model_params_lowvram_keys[f"{n}.weight"] = n
|
||||
|
||||
def unpatch_model(self, device_to=None, unpatch_weights=True, *args, **kwargs):
|
||||
try:
|
||||
return super().unpatch_model(device_to, unpatch_weights, *args, **kwargs)
|
||||
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:
|
||||
comfy.utils.copy_to_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1]))
|
||||
else:
|
||||
for k in keys:
|
||||
comfy.utils.set_attr_param(self.model, k, self.hooked_backup[k][0].to(device=self.hooked_backup[k][1]))
|
||||
self.model_params_lowvram = False
|
||||
self.model_params_lowvram_keys.clear()
|
||||
# clear hooked_backup
|
||||
self.hooked_backup.clear()
|
||||
self.current_lora_hooks = None
|
||||
|
||||
|
||||
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,
|
||||
strength_model: float, strength_clip: float):
|
||||
if model is not None and model_loaded is not None:
|
||||
new_modelpatcher = ModelPatcherAndInjector.create_from(model)
|
||||
comfy.model_management.unload_model_clones(new_modelpatcher)
|
||||
expected_model_keys = model_loaded.model_keys.copy()
|
||||
patches_model: dict[str, Tensor] = model_loaded.model.state_dict()
|
||||
# do not include ANY model_sampling components of the model that should act as a patch
|
||||
for key in list(patches_model.keys()):
|
||||
if key.startswith("model_sampling"):
|
||||
expected_model_keys.discard(key)
|
||||
patches_model.pop(key, None)
|
||||
k = new_modelpatcher.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=patches_model, strength_patch=strength_model)
|
||||
else:
|
||||
k = ()
|
||||
new_modelpatcher = None
|
||||
|
||||
if clip is not None and clip_loaded is not None:
|
||||
new_clip = CLIPWithHooks(clip)
|
||||
comfy.model_management.unload_model_clones(new_clip.patcher)
|
||||
expected_clip_keys = clip_loaded.patcher.model_keys.copy()
|
||||
patches_clip: dict[str, Tensor] = clip_loaded.cond_stage_model.state_dict()
|
||||
k1 = new_clip.add_hooked_patches_as_diffs(lora_hook=lora_hook, patches=patches_clip, strength_patch=strength_clip)
|
||||
else:
|
||||
k1 = ()
|
||||
new_clip = None
|
||||
|
||||
k = set(k)
|
||||
k1 = set(k1)
|
||||
if model is not None and model_loaded is not None:
|
||||
for key in expected_model_keys:
|
||||
if key not in k:
|
||||
logger.warning(f"MODEL-AS-LORA NOT LOADED {key}")
|
||||
if clip is not None and clip_loaded is not None:
|
||||
for key in expected_clip_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
|
||||
@@ -396,6 +964,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
|
||||
|
||||
+59
-9
@@ -10,6 +10,14 @@ from .nodes_cameractrl import (LoadAnimateDiffModelWithCameraCtrl, ApplyAnimateD
|
||||
CameraCtrlPoseBasic, CameraCtrlPoseCombo, CameraCtrlPoseAdvanced, CameraCtrlManualAppendPose,
|
||||
CameraCtrlReplaceCameraParameters, CameraCtrlSetOriginalAspectRatio)
|
||||
from .nodes_multival import MultivalDynamicNode, MultivalScaledMaskNode
|
||||
from .nodes_conditioning import (MaskableLoraLoader, MaskableLoraLoaderModelOnly, MaskableSDModelLoader, MaskableSDModelLoaderModelOnly,
|
||||
SetModelLoraHook, SetClipLoraHook,
|
||||
CombineLoraHooks, CombineLoraHookFourOptional, CombineLoraHookEightOptional,
|
||||
PairedConditioningSetMaskHooked, ConditioningSetMaskHooked,
|
||||
PairedConditioningSetMaskAndCombineHooked, ConditioningSetMaskAndCombineHooked,
|
||||
PairedConditioningSetUnmaskedAndCombineHooked, ConditioningSetUnmaskedAndCombineHooked,
|
||||
ConditioningTimestepsNode, SetLoraHookKeyframes,
|
||||
CreateLoraHookKeyframe, CreateLoraHookKeyframeInterpolation, CreateLoraHookKeyframeFromStrengthList)
|
||||
from .nodes_sample import (FreeInitOptionsNode, NoiseLayerAddWeightedNode, SampleSettingsNode, NoiseLayerAddNode, NoiseLayerReplaceNode, IterationOptionsNode,
|
||||
CustomCFGNode, CustomCFGKeyframeNode)
|
||||
from .nodes_sigma_schedule import (SigmaScheduleNode, RawSigmaScheduleNode, WeightedAverageSigmaScheduleNode, InterpolatedWeightedAverageSigmaScheduleNode, SplitAndCombineSigmaScheduleNode)
|
||||
@@ -21,7 +29,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
|
||||
|
||||
@@ -52,6 +60,27 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Iteration Opts
|
||||
"ADE_IterationOptsDefault": IterationOptionsNode,
|
||||
"ADE_IterationOptsFreeInit": FreeInitOptionsNode,
|
||||
# Conditioning
|
||||
"ADE_RegisterLoraHook": MaskableLoraLoader,
|
||||
"ADE_RegisterLoraHookModelOnly": MaskableLoraLoaderModelOnly,
|
||||
"ADE_RegisterModelAsLoraHook": MaskableSDModelLoader,
|
||||
"ADE_RegisterModelAsLoraHookModelOnly": MaskableSDModelLoaderModelOnly,
|
||||
"ADE_CombineLoraHooks": CombineLoraHooks,
|
||||
"ADE_CombineLoraHooksFour": CombineLoraHookFourOptional,
|
||||
"ADE_CombineLoraHooksEight": CombineLoraHookEightOptional,
|
||||
"ADE_SetLoraHookKeyframe": SetLoraHookKeyframes,
|
||||
"ADE_AttachLoraHookToCLIP": SetClipLoraHook,
|
||||
"ADE_LoraHookKeyframe": CreateLoraHookKeyframe,
|
||||
"ADE_LoraHookKeyframeInterpolation": CreateLoraHookKeyframeInterpolation,
|
||||
"ADE_LoraHookKeyframeFromStrengthList": CreateLoraHookKeyframeFromStrengthList,
|
||||
"ADE_AttachLoraHookToConditioning": SetModelLoraHook,
|
||||
"ADE_PairedConditioningSetMask": PairedConditioningSetMaskHooked,
|
||||
"ADE_ConditioningSetMask": ConditioningSetMaskHooked,
|
||||
"ADE_PairedConditioningSetMaskAndCombine": PairedConditioningSetMaskAndCombineHooked,
|
||||
"ADE_ConditioningSetMaskAndCombine": ConditioningSetMaskAndCombineHooked,
|
||||
"ADE_PairedConditioningSetUnmaskedAndCombine": PairedConditioningSetUnmaskedAndCombineHooked,
|
||||
"ADE_ConditioningSetUnmaskedAndCombine": ConditioningSetUnmaskedAndCombineHooked,
|
||||
"ADE_TimestepsConditioning": ConditioningTimestepsNode,
|
||||
# Noise Layer Nodes
|
||||
"ADE_NoiseLayerAdd": NoiseLayerAddNode,
|
||||
"ADE_NoiseLayerAddWeighted": NoiseLayerAddWeightedNode,
|
||||
@@ -82,10 +111,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,
|
||||
@@ -113,6 +138,10 @@ NODE_CLASS_MAPPINGS = {
|
||||
"AnimateDiffLoaderV1": AnimateDiffLoader_Deprecated,
|
||||
"ADE_AnimateDiffLoaderV1Advanced": AnimateDiffLoaderAdvanced_Deprecated,
|
||||
"ADE_AnimateDiffCombine": AnimateDiffCombine_Deprecated,
|
||||
"ADE_AnimateDiffModelSettings_Release": AnimateDiffModelSettings,
|
||||
"ADE_AnimateDiffModelSettingsSimple": AnimateDiffModelSettingsSimple,
|
||||
"ADE_AnimateDiffModelSettings": AnimateDiffModelSettingsAdvanced,
|
||||
"ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": AnimateDiffModelSettingsAdvancedAttnStrengths,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Unencapsulated
|
||||
@@ -136,6 +165,27 @@ 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 🎭🅐🅓",
|
||||
"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_SetLoraHookKeyframe": "Set LoRA Hook Keyframes 🎭🅐🅓",
|
||||
"ADE_AttachLoraHookToCLIP": "Set CLIP LoRA Hook 🎭🅐🅓",
|
||||
"ADE_LoraHookKeyframe": "LoRA Hook Keyframe 🎭🅐🅓",
|
||||
"ADE_LoraHookKeyframeInterpolation": "LoRA Hook Keyframes Interpolation 🎭🅐🅓",
|
||||
"ADE_LoraHookKeyframeFromStrengthList": "LoRA Hook Keyframes From List 🎭🅐🅓",
|
||||
"ADE_AttachLoraHookToConditioning": "Set Model LoRA Hook 🎭🅐🅓",
|
||||
"ADE_PairedConditioningSetMask": "Set Props on Conds 🎭🅐🅓",
|
||||
"ADE_ConditioningSetMask": "Set Props on Cond 🎭🅐🅓",
|
||||
"ADE_PairedConditioningSetMaskAndCombine": "Set Props and Combine Conds 🎭🅐🅓",
|
||||
"ADE_ConditioningSetMaskAndCombine": "Set Props and Combine Cond 🎭🅐🅓",
|
||||
"ADE_PairedConditioningSetUnmaskedAndCombine": "Set Unmasked Conds 🎭🅐🅓",
|
||||
"ADE_ConditioningSetUnmaskedAndCombine": "Set Unmasked Cond 🎭🅐🅓",
|
||||
"ADE_TimestepsConditioning": "Timesteps Conditioning 🎭🅐🅓",
|
||||
# Noise Layer Nodes
|
||||
"ADE_NoiseLayerAdd": "Noise Layer [Add] 🎭🅐🅓",
|
||||
"ADE_NoiseLayerAddWeighted": "Noise Layer [Add Weighted] 🎭🅐🅓",
|
||||
@@ -166,10 +216,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 🎭🅐🅓②",
|
||||
@@ -197,4 +243,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AnimateDiffLoaderV1": "🚫AnimateDiff Loader [DEPRECATED] 🎭🅐🅓",
|
||||
"ADE_AnimateDiffLoaderV1Advanced": "🚫AnimateDiff Loader (Advanced) [DEPRECATED] 🎭🅐🅓",
|
||||
"ADE_AnimateDiffCombine": "🚫AnimateDiff Combine [DEPRECATED, Use Video Combine (VHS) Instead!] 🎭🅐🅓",
|
||||
"ADE_AnimateDiffModelSettings_Release": "🚫[DEPR] Motion Model Settings 🎭🅐🅓①",
|
||||
"ADE_AnimateDiffModelSettingsSimple": "🚫[DEPR] Motion Model Settings (Simple) 🎭🅐🅓①",
|
||||
"ADE_AnimateDiffModelSettings": "🚫[DEPR] Motion Model Settings (Advanced) 🎭🅐🅓①",
|
||||
"ADE_AnimateDiffModelSettingsAdvancedAttnStrengths": "🚫[DEPR] Motion Model Settings (Adv. Attn) 🎭🅐🅓①",
|
||||
}
|
||||
|
||||
@@ -0,0 +1,624 @@
|
||||
import uuid
|
||||
import folder_paths
|
||||
from typing import Union
|
||||
from torch import Tensor
|
||||
from collections.abc import Iterable
|
||||
|
||||
from comfy.model_patcher import ModelPatcher
|
||||
from comfy.sd import CLIP
|
||||
import comfy.sd
|
||||
import comfy.utils
|
||||
|
||||
from .conditioning import (COND_CONST, TimestepsCond, set_mask_conds, set_mask_and_combine_conds, set_unmasked_and_combine_conds,
|
||||
LoraHook, LoraHookGroup, LoraHookKeyframe, LoraHookKeyframeGroup)
|
||||
from .model_injection import ModelPatcherAndInjector, CLIPWithHooks, load_hooked_lora_for_models, load_model_as_hooked_lora_for_models
|
||||
from .utils_model import BIGMAX, InterpolationMethod
|
||||
from .logger import logger
|
||||
|
||||
|
||||
###############################################
|
||||
### Mask, Combine, and Hook Conditioning
|
||||
###############################################
|
||||
class PairedConditioningSetMaskHooked:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"positive_ADD": ("CONDITIONING", ),
|
||||
"negative_ADD": ("CONDITIONING", ),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"set_cond_area": (COND_CONST._LIST_COND_AREA,),
|
||||
},
|
||||
"optional": {
|
||||
"opt_mask": ("MASK", ),
|
||||
"opt_lora_hook": ("LORA_HOOK",),
|
||||
"opt_timesteps": ("TIMESTEPS_COND",)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
|
||||
RETURN_NAMES = ("positive", "negative")
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning"
|
||||
FUNCTION = "append_and_hook"
|
||||
|
||||
def append_and_hook(self, positive_ADD, negative_ADD,
|
||||
strength: float, set_cond_area: str,
|
||||
opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None):
|
||||
final_positive, final_negative = set_mask_conds(conds=[positive_ADD, negative_ADD],
|
||||
strength=strength, set_cond_area=set_cond_area,
|
||||
opt_mask=opt_mask, opt_lora_hook=opt_lora_hook, opt_timesteps=opt_timesteps)
|
||||
return (final_positive, final_negative)
|
||||
|
||||
|
||||
class ConditioningSetMaskHooked:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"cond_ADD": ("CONDITIONING",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"set_cond_area": (COND_CONST._LIST_COND_AREA,),
|
||||
},
|
||||
"optional": {
|
||||
"opt_mask": ("MASK", ),
|
||||
"opt_lora_hook": ("LORA_HOOK",),
|
||||
"opt_timesteps": ("TIMESTEPS_COND",)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/single cond ops"
|
||||
FUNCTION = "append_and_hook"
|
||||
|
||||
def append_and_hook(self, cond_ADD,
|
||||
strength: float, set_cond_area: str,
|
||||
opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None):
|
||||
(final_conditioning,) = set_mask_conds(conds=[cond_ADD],
|
||||
strength=strength, set_cond_area=set_cond_area,
|
||||
opt_mask=opt_mask, opt_lora_hook=opt_lora_hook, opt_timesteps=opt_timesteps)
|
||||
return (final_conditioning,)
|
||||
|
||||
|
||||
class PairedConditioningSetMaskAndCombineHooked:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
"positive_ADD": ("CONDITIONING",),
|
||||
"negative_ADD": ("CONDITIONING",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"set_cond_area": (COND_CONST._LIST_COND_AREA,),
|
||||
},
|
||||
"optional": {
|
||||
"opt_mask": ("MASK", ),
|
||||
"opt_lora_hook": ("LORA_HOOK",),
|
||||
"opt_timesteps": ("TIMESTEPS_COND",)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
|
||||
RETURN_NAMES = ("positive", "negative")
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning"
|
||||
FUNCTION = "append_and_combine"
|
||||
|
||||
def append_and_combine(self, positive, negative, positive_ADD, negative_ADD,
|
||||
strength: float, set_cond_area: str,
|
||||
opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None):
|
||||
final_positive, final_negative = set_mask_and_combine_conds(conds=[positive, negative], new_conds=[positive_ADD, negative_ADD],
|
||||
strength=strength, set_cond_area=set_cond_area,
|
||||
opt_mask=opt_mask, opt_lora_hook=opt_lora_hook, opt_timesteps=opt_timesteps)
|
||||
return (final_positive, final_negative,)
|
||||
|
||||
|
||||
class ConditioningSetMaskAndCombineHooked:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"cond": ("CONDITIONING",),
|
||||
"cond_ADD": ("CONDITIONING",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"set_cond_area": (COND_CONST._LIST_COND_AREA,),
|
||||
},
|
||||
"optional": {
|
||||
"opt_mask": ("MASK", ),
|
||||
"opt_lora_hook": ("LORA_HOOK",),
|
||||
"opt_timesteps": ("TIMESTEPS_COND",)
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/single cond ops"
|
||||
FUNCTION = "append_and_combine"
|
||||
|
||||
def append_and_combine(self, conditioning, conditioning_ADD,
|
||||
strength: float, set_cond_area: str,
|
||||
opt_mask: Tensor=None, opt_lora_hook: LoraHookGroup=None, opt_timesteps: TimestepsCond=None):
|
||||
(final_conditioning,) = set_mask_and_combine_conds(conds=[conditioning], new_conds=[conditioning_ADD],
|
||||
strength=strength, set_cond_area=set_cond_area,
|
||||
opt_mask=opt_mask, opt_lora_hook=opt_lora_hook, opt_timesteps=opt_timesteps)
|
||||
return (final_conditioning,)
|
||||
|
||||
|
||||
class PairedConditioningSetUnmaskedAndCombineHooked:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"positive": ("CONDITIONING",),
|
||||
"negative": ("CONDITIONING",),
|
||||
"positive_DEFAULT": ("CONDITIONING",),
|
||||
"negative_DEFAULT": ("CONDITIONING",),
|
||||
},
|
||||
"optional": {
|
||||
"opt_lora_hook": ("LORA_HOOK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
|
||||
RETURN_NAMES = ("positive", "negative")
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning"
|
||||
FUNCTION = "append_and_combine"
|
||||
|
||||
def append_and_combine(self, positive, negative, positive_DEFAULT, negative_DEFAULT,
|
||||
opt_lora_hook: LoraHookGroup=None):
|
||||
final_positive, final_negative = set_unmasked_and_combine_conds(conds=[positive, negative], new_conds=[positive_DEFAULT, negative_DEFAULT],
|
||||
opt_lora_hook=opt_lora_hook)
|
||||
return (final_positive, final_negative,)
|
||||
|
||||
|
||||
class ConditioningSetUnmaskedAndCombineHooked:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"cond": ("CONDITIONING",),
|
||||
"cond_DEFAULT": ("CONDITIONING",),
|
||||
},
|
||||
"optional": {
|
||||
"opt_lora_hook": ("LORA_HOOK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/single cond ops"
|
||||
FUNCTION = "append_and_combine"
|
||||
|
||||
def append_and_combine(self, cond, cond_DEFAULT,
|
||||
opt_lora_hook: LoraHookGroup=None):
|
||||
(final_conditioning,) = set_unmasked_and_combine_conds(conds=[cond], new_conds=[cond_DEFAULT],
|
||||
opt_lora_hook=opt_lora_hook)
|
||||
return (final_conditioning,)
|
||||
###############################################
|
||||
###############################################
|
||||
###############################################
|
||||
|
||||
|
||||
|
||||
###############################################
|
||||
### Scheduling
|
||||
###############################################
|
||||
class ConditioningTimestepsNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TIMESTEPS_COND",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning"
|
||||
FUNCTION = "create_schedule"
|
||||
|
||||
def create_schedule(self, start_percent: float, end_percent: float):
|
||||
return (TimestepsCond(start_percent=start_percent, end_percent=end_percent),)
|
||||
|
||||
|
||||
class SetLoraHookKeyframes:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"lora_hook": ("LORA_HOOK",),
|
||||
"hook_kf": ("LORA_HOOK_KEYFRAMES",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LORA_HOOK",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning"
|
||||
FUNCTION = "set_hook_keyframes"
|
||||
|
||||
def set_hook_keyframes(self, lora_hook: LoraHookGroup, hook_kf: LoraHookKeyframeGroup):
|
||||
new_lora_hook = lora_hook.clone()
|
||||
new_lora_hook.set_keyframes_on_hooks(hook_kf=hook_kf)
|
||||
return (new_lora_hook,)
|
||||
|
||||
|
||||
class CreateLoraHookKeyframe:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"guarantee_steps": ("INT", {"default": 1, "min": 0, "max": BIGMAX}),
|
||||
},
|
||||
"optional": {
|
||||
"prev_hook_kf": ("LORA_HOOK_KEYFRAMES",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LORA_HOOK_KEYFRAMES",)
|
||||
RETURN_NAMES = ("HOOK_KF",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/schedule lora hooks"
|
||||
FUNCTION = "create_hook_keyframe"
|
||||
|
||||
def create_hook_keyframe(self, strength_model: float, start_percent: float, guarantee_steps: float,
|
||||
prev_hook_kf: LoraHookKeyframeGroup=None):
|
||||
if prev_hook_kf:
|
||||
prev_hook_kf = prev_hook_kf.clone()
|
||||
else:
|
||||
prev_hook_kf = LoraHookKeyframeGroup()
|
||||
keyframe = LoraHookKeyframe(strength=strength_model, start_percent=start_percent, guarantee_steps=guarantee_steps)
|
||||
prev_hook_kf.add(keyframe)
|
||||
return (prev_hook_kf,)
|
||||
|
||||
|
||||
class CreateLoraHookKeyframeInterpolation:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"strength_start": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"strength_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||
"interpolation": (InterpolationMethod._LIST, ),
|
||||
"intervals": ("INT", {"default": 5, "min": 2, "max": 100, "step": 1}),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"prev_hook_kf": ("LORA_HOOK_KEYFRAMES",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LORA_HOOK_KEYFRAMES",)
|
||||
RETURN_NAMES = ("HOOK_KF",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/schedule lora hooks"
|
||||
FUNCTION = "create_hook_keyframes"
|
||||
|
||||
def create_hook_keyframes(self,
|
||||
start_percent: float, end_percent: float,
|
||||
strength_start: float, strength_end: float, interpolation: str, intervals: int,
|
||||
prev_hook_kf: LoraHookKeyframeGroup=None, print_keyframes=False):
|
||||
if prev_hook_kf:
|
||||
prev_hook_kf = prev_hook_kf.clone()
|
||||
else:
|
||||
prev_hook_kf = LoraHookKeyframeGroup()
|
||||
percents = InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=intervals, method=interpolation)
|
||||
strengths = InterpolationMethod.get_weights(num_from=strength_start, num_to=strength_end, length=intervals, method=interpolation)
|
||||
|
||||
is_first = True
|
||||
for percent, strength in zip(percents, strengths):
|
||||
guarantee_steps = 0
|
||||
if is_first:
|
||||
guarantee_steps = 1
|
||||
is_first = False
|
||||
prev_hook_kf.add(LoraHookKeyframe(strength=strength, start_percent=percent, guarantee_steps=guarantee_steps))
|
||||
if print_keyframes:
|
||||
logger.info(f"LoraHookKeyframe - start_percent:{percent} = {strength}")
|
||||
return (prev_hook_kf,)
|
||||
|
||||
|
||||
class CreateLoraHookKeyframeFromStrengthList:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"strengths_float": ("FLOAT", {"default": -1, "min": -1, "step": 0.001, "forceInput": True}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||
},
|
||||
"optional": {
|
||||
"prev_hook_kf": ("LORA_HOOK_KEYFRAMES",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LORA_HOOK_KEYFRAMES",)
|
||||
RETURN_NAMES = ("HOOK_KF",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/schedule lora hooks"
|
||||
FUNCTION = "create_hook_keyframes"
|
||||
|
||||
def create_hook_keyframes(self, strengths_float: Union[float, list[float]],
|
||||
start_percent: float, end_percent: float,
|
||||
prev_hook_kf: LoraHookKeyframeGroup=None, print_keyframes=False):
|
||||
if prev_hook_kf:
|
||||
prev_hook_kf = prev_hook_kf.clone()
|
||||
else:
|
||||
prev_hook_kf = LoraHookKeyframeGroup()
|
||||
if type(strengths_float) in (float, int):
|
||||
strengths_float = [float(strengths_float)]
|
||||
elif isinstance(strengths_float, Iterable):
|
||||
pass
|
||||
else:
|
||||
raise Exception(f"strengths_floast must be either an interable input or a float, but was {type(strengths_float).__repr__}.")
|
||||
percents = InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=len(strengths_float), method=InterpolationMethod.LINEAR)
|
||||
|
||||
is_first = True
|
||||
for percent, strength in zip(percents, strengths_float):
|
||||
guarantee_steps = 0
|
||||
if is_first:
|
||||
guarantee_steps = 1
|
||||
is_first = False
|
||||
prev_hook_kf.add(LoraHookKeyframe(strength=strength, start_percent=percent, guarantee_steps=guarantee_steps))
|
||||
if print_keyframes:
|
||||
logger.info(f"LoraHookKeyframe - start_percent:{percent} = {strength}")
|
||||
return (prev_hook_kf,)
|
||||
###############################################
|
||||
###############################################
|
||||
###############################################
|
||||
|
||||
|
||||
###############################################
|
||||
### Register LoRA Hooks
|
||||
###############################################
|
||||
# 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/register lora hooks"
|
||||
FUNCTION = "load_lora"
|
||||
|
||||
def load_lora(self, model: Union[ModelPatcher, ModelPatcherAndInjector], clip: CLIP, lora_name: str, strength_model: float, strength_clip: float):
|
||||
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/register lora hooks"
|
||||
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"), ),
|
||||
"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/register lora hooks"
|
||||
FUNCTION = "load_model_as_lora"
|
||||
|
||||
def load_model_as_lora(self, model: ModelPatcher, clip: CLIP, ckpt_name: str, strength_model: float, strength_clip: float):
|
||||
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||
out = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings"))
|
||||
model_loaded = out[0]
|
||||
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,
|
||||
strength_model=strength_model, strength_clip=strength_clip)
|
||||
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"), ),
|
||||
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "LORA_HOOK")
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/register lora hooks"
|
||||
FUNCTION = "load_model_as_lora_model_only"
|
||||
|
||||
def load_model_as_lora_model_only(self, model: ModelPatcher, ckpt_name: str, strength_model: float):
|
||||
model_lora, clip_lora, lora_hook = self.load_model_as_lora(model=model, clip=None, ckpt_name=ckpt_name,
|
||||
strength_model=strength_model, strength_clip=0)
|
||||
return (model_lora, lora_hook)
|
||||
###############################################
|
||||
###############################################
|
||||
###############################################
|
||||
|
||||
|
||||
|
||||
###############################################
|
||||
### Set LoRA Hooks
|
||||
###############################################
|
||||
class SetModelLoraHook:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"conditioning": ("CONDITIONING",),
|
||||
"lora_hook": ("LORA_HOOK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/single cond ops"
|
||||
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": {
|
||||
},
|
||||
"optional": {
|
||||
"lora_hook_A": ("LORA_HOOK",),
|
||||
"lora_hook_B": ("LORA_HOOK",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LORA_HOOK",)
|
||||
CATEGORY = "Animate Diff 🎭🅐🅓/conditioning/combine lora hooks"
|
||||
FUNCTION = "combine_lora_hooks"
|
||||
|
||||
def combine_lora_hooks(self, lora_hook_A: LoraHookGroup=None, lora_hook_B: LoraHookGroup=None):
|
||||
candidates = [lora_hook_A, lora_hook_B]
|
||||
return (LoraHookGroup.combine_all_lora_hooks(candidates),)
|
||||
|
||||
|
||||
class CombineLoraHookFourOptional:
|
||||
@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 lora 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 lora 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
|
||||
###############################################
|
||||
###############################################
|
||||
###############################################
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -52,7 +52,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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ from comfy.model_patcher import ModelPatcher
|
||||
from comfy.model_base import BaseModel
|
||||
|
||||
from . import freeinit
|
||||
from .conditioning import LoraHookMode
|
||||
from .context import ContextOptions, ContextOptionsGroup
|
||||
from .utils_model import SigmaSchedule
|
||||
from .utils_motion import extend_to_batch_size, get_sorted_list_via_attr, prepare_mask_batch
|
||||
|
||||
+360
-67
@@ -1,5 +1,6 @@
|
||||
from typing import Callable
|
||||
|
||||
import collections
|
||||
import math
|
||||
import torch
|
||||
from torch import Tensor
|
||||
@@ -18,8 +19,10 @@ except ImportError:
|
||||
SAMPLE_FALLBACK = True
|
||||
import comfy.utils
|
||||
from comfy.controlnet import ControlBase
|
||||
from comfy.model_base import BaseModel
|
||||
import comfy.ops
|
||||
|
||||
from .conditioning import COND_CONST, LoraHookGroup
|
||||
from .context import ContextFuseMethod, ContextSchedules, get_context_weights, get_context_windows
|
||||
from .sample_settings import IterationOptions, SampleSettings, SeedNoiseGeneration
|
||||
from .utils_model import ModelTypeSD
|
||||
@@ -33,12 +36,13 @@ 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
|
||||
self.reset()
|
||||
|
||||
def initialize(self, model):
|
||||
|
||||
def initialize(self, model: BaseModel):
|
||||
# this function is to be run in sampling func
|
||||
if not self.initialized:
|
||||
self.initialized = True
|
||||
@@ -49,12 +53,38 @@ class AnimateDiffHelper_GlobalState:
|
||||
if self.sample_settings.custom_cfg is not None:
|
||||
self.sample_settings.custom_cfg.initialize_timesteps(model)
|
||||
|
||||
def hooks_initialize(self, model: BaseModel, hook_groups: list[LoraHookGroup]):
|
||||
# this function is to be run the first time all gathered
|
||||
if not self.hooks_initialized:
|
||||
self.hooks_initialized = True
|
||||
for hook_group in hook_groups:
|
||||
for hook in hook_group.hooks:
|
||||
hook.reset()
|
||||
hook.initialize_timesteps(model)
|
||||
|
||||
def prepare_current_keyframes(self, timestep: Tensor):
|
||||
if self.motion_models is not None:
|
||||
self.motion_models.prepare_current_keyframe(t=timestep)
|
||||
if self.params.context_options is not None:
|
||||
self.params.context_options.prepare_current_context(t=timestep)
|
||||
if self.sample_settings.custom_cfg is not None:
|
||||
self.sample_settings.custom_cfg.prepare_current_keyframe(t=timestep)
|
||||
|
||||
def prepare_hooks_current_keyframes(self, timestep: Tensor, hook_groups: list[LoraHookGroup]):
|
||||
if self.model_patcher is not None:
|
||||
self.model_patcher.prepare_hooked_patches_current_keyframe(t=timestep, hook_groups=hook_groups)
|
||||
|
||||
def reset(self):
|
||||
self.initialized = False
|
||||
self.hooks_initialized = False
|
||||
self.start_step: int = 0
|
||||
self.last_step: int = 0
|
||||
self.current_step: int = 0
|
||||
self.total_steps: int = 0
|
||||
if self.model_patcher is not None:
|
||||
self.model_patcher.clean_hooks()
|
||||
del self.model_patcher
|
||||
self.model_patcher = None
|
||||
if self.motion_models is not None:
|
||||
del self.motion_models
|
||||
self.motion_models = None
|
||||
@@ -64,13 +94,13 @@ class AnimateDiffHelper_GlobalState:
|
||||
if self.sample_settings is not None:
|
||||
del self.sample_settings
|
||||
self.sample_settings = None
|
||||
|
||||
|
||||
def update_with_inject_params(self, params: InjectionParams):
|
||||
self.params = params
|
||||
|
||||
def is_using_sliding_context(self):
|
||||
return self.params is not None and self.params.is_using_sliding_context()
|
||||
|
||||
|
||||
def create_exposed_params(self):
|
||||
# This dict will be exposed to be used by other extensions
|
||||
# DO NOT change any of the key names
|
||||
@@ -214,12 +244,13 @@ class FunctionInjectionHolder:
|
||||
pass
|
||||
|
||||
def inject_functions(self, model: ModelPatcherAndInjector, params: InjectionParams):
|
||||
# Save Original Functions
|
||||
# Save Original Functions - order must match between here and restore_functions
|
||||
self.orig_forward_timestep_embed = openaimodel.forward_timestep_embed # needed to account for VanillaTemporalModule
|
||||
self.orig_memory_required = model.model.memory_required # allows for "unlimited area hack" to prevent halving of conds/unconds
|
||||
self.orig_groupnorm_forward = torch.nn.GroupNorm.forward # used to normalize latents to remove "flickering" of colors/brightness between frames
|
||||
self.orig_groupnorm_manual_cast_forward = comfy.ops.manual_cast.GroupNorm.forward_comfy_cast_weights
|
||||
self.orig_sampling_function = comfy.samplers.sampling_function # used to support sliding context windows in samplers
|
||||
self.orig_get_area_and_mult = comfy.samplers.get_area_and_mult
|
||||
if SAMPLE_FALLBACK: # for backwards compatibility, for now
|
||||
self.orig_get_additional_models = comfy.sample.get_additional_models
|
||||
else:
|
||||
@@ -249,6 +280,7 @@ class FunctionInjectionHolder:
|
||||
break
|
||||
del info
|
||||
comfy.samplers.sampling_function = evolved_sampling_function
|
||||
comfy.samplers.get_area_and_mult = get_area_and_mult_ADE
|
||||
if SAMPLE_FALLBACK: # for backwards compatibility, for now
|
||||
comfy.sample.get_additional_models = get_additional_models_factory(self.orig_get_additional_models, model.motion_models)
|
||||
else:
|
||||
@@ -261,6 +293,7 @@ class FunctionInjectionHolder:
|
||||
openaimodel.forward_timestep_embed = self.orig_forward_timestep_embed
|
||||
torch.nn.GroupNorm.forward = self.orig_groupnorm_forward
|
||||
comfy.ops.manual_cast.GroupNorm.forward_comfy_cast_weights = self.orig_groupnorm_manual_cast_forward
|
||||
comfy.samplers.get_area_and_mult = self.orig_get_area_and_mult
|
||||
comfy.samplers.sampling_function = self.orig_sampling_function
|
||||
if SAMPLE_FALLBACK: # for backwards compatibility, for now
|
||||
comfy.sample.get_additional_models = self.orig_get_additional_models
|
||||
@@ -318,6 +351,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
|
||||
|
||||
@@ -366,6 +400,7 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) ->
|
||||
ADGS.start_step = kwargs.get("start_step") or 0
|
||||
ADGS.current_step = ADGS.start_step
|
||||
ADGS.last_step = kwargs.get("last_step") or 0
|
||||
ADGS.hooks_initialized = False
|
||||
if iter_opts.iterations > 1:
|
||||
logger.info(f"Iteration {curr_i+1}/{iter_opts.iterations}")
|
||||
# perform any iter_opts preprocessing on latents
|
||||
@@ -400,12 +435,7 @@ def motion_sample_factory(orig_comfy_sample: Callable, is_custom: bool=False) ->
|
||||
|
||||
def evolved_sampling_function(model, x, timestep, uncond, cond, cond_scale, model_options: dict={}, seed=None):
|
||||
ADGS.initialize(model)
|
||||
if ADGS.motion_models is not None:
|
||||
ADGS.motion_models.prepare_current_keyframe(t=timestep)
|
||||
if ADGS.params.context_options is not None:
|
||||
ADGS.params.context_options.prepare_current_context(t=timestep)
|
||||
if ADGS.sample_settings.custom_cfg is not None:
|
||||
ADGS.sample_settings.custom_cfg.prepare_current_keyframe(t=timestep)
|
||||
ADGS.prepare_current_keyframes(timestep=timestep)
|
||||
|
||||
# never use cfg1 optimization if using custom_cfg (since can have timesteps and such)
|
||||
if ADGS.sample_settings.custom_cfg is None and math.isclose(cond_scale, 1.0) and model_options.get("disable_cfg1_optimization", False) == False:
|
||||
@@ -420,19 +450,15 @@ 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)
|
||||
cond_pred, uncond_pred = sliding_calc_conds_batch(model, [cond, uncond_], x, timestep, model_options)
|
||||
|
||||
if hasattr(comfy.samplers, "cfg_function"):
|
||||
try:
|
||||
cached_calc_cond_batch = comfy.samplers.calc_cond_batch
|
||||
# support sliding context for PAG/other sampler_post_cfg_function tech that may use calc_cond_batch
|
||||
if ADGS.is_using_sliding_context():
|
||||
comfy.samplers.calc_cond_batch = wrapped_cfg_sliding_calc_cond_batch_factory(cached_calc_cond_batch)
|
||||
# support hooks and sliding context for PAG/other sampler_post_cfg_function tech that may use calc_cond_batch
|
||||
comfy.samplers.calc_cond_batch = wrapped_cfg_sliding_calc_cond_batch_factory(cached_calc_cond_batch)
|
||||
return comfy.samplers.cfg_function(model, cond_pred, uncond_pred, cond_scale, x, timestep, model_options, cond, uncond)
|
||||
finally:
|
||||
comfy.samplers.calc_cond_batch = cached_calc_cond_batch
|
||||
@@ -456,34 +482,35 @@ def wrapped_cfg_sliding_calc_cond_batch_factory(orig_calc_cond_batch):
|
||||
def wrapped_cfg_sliding_calc_cond_batch(model, conds, x_in, timestep, model_options):
|
||||
# current call to calc_cond_batch should refer to sliding version
|
||||
try:
|
||||
uncond = None
|
||||
current_calc_cond_batch = comfy.samplers.calc_cond_batch
|
||||
# when inside sliding_calc_cond_uncond, should return to original calc_cond_batch
|
||||
# when inside sliding_calc_conds_batch, should return to original calc_cond_batch
|
||||
comfy.samplers.calc_cond_batch = orig_calc_cond_batch
|
||||
if len(conds) > 1:
|
||||
uncond = conds[1]
|
||||
result = sliding_calc_cond_uncond_batch(model, conds[0], uncond, x_in, timestep, model_options)
|
||||
if uncond is None:
|
||||
result = (result[0],)
|
||||
return result
|
||||
if not ADGS.is_using_sliding_context():
|
||||
return calc_cond_uncond_batch_wrapper(model, conds, x_in, timestep, model_options)
|
||||
else:
|
||||
return sliding_calc_conds_batch(model, conds, x_in, timestep, model_options)
|
||||
finally:
|
||||
del uncond
|
||||
# make sure calc_cond_batch will become wrapped again
|
||||
comfy.samplers.calc_cond_batch = current_calc_cond_batch
|
||||
return wrapped_cfg_sliding_calc_cond_batch
|
||||
|
||||
|
||||
# sliding_calc_cond_uncond_batch inspired by ashen's initial hack for 16-frame sliding context:
|
||||
# sliding_calc_conds_batch inspired by ashen's initial hack for 16-frame sliding context:
|
||||
# https://github.com/comfyanonymous/ComfyUI/compare/master...ashen-sensored:ComfyUI:master
|
||||
def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep, model_options):
|
||||
def sliding_calc_conds_batch(model, conds, x_in: Tensor, timestep, model_options):
|
||||
def prepare_control_objects(control: ControlBase, full_idxs: list[int]):
|
||||
if control.previous_controlnet is not None:
|
||||
prepare_control_objects(control.previous_controlnet, full_idxs)
|
||||
if not hasattr(control, "sub_idxs"):
|
||||
raise ValueError(f"Control type {type(control).__name__} may not support required features for sliding context window; \
|
||||
use ControlNet nodes from Kosinkadink/ComfyUI-Advanced-ControlNet, or make sure ComfyUI-Advanced-ControlNet is updated.")
|
||||
control.sub_idxs = full_idxs
|
||||
control.full_latent_length = ADGS.params.full_length
|
||||
control.context_length = ADGS.params.context_options.context_length
|
||||
|
||||
def get_resized_cond(cond_in, full_idxs: list[int], context_length: int) -> list:
|
||||
if cond_in is None:
|
||||
return None
|
||||
# reuse or resize cond items to match context requirements
|
||||
resized_cond = []
|
||||
# cond object is a list containing a dict - outer list is irrelevant, so just loop through it
|
||||
@@ -504,11 +531,7 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep,
|
||||
# look for control
|
||||
elif key == "control":
|
||||
control_item = cond_item
|
||||
if hasattr(control_item, "sub_idxs"):
|
||||
prepare_control_objects(control_item, full_idxs)
|
||||
else:
|
||||
raise ValueError(f"Control type {type(control_item).__name__} may not support required features for sliding context window; \
|
||||
use Control objects from Kosinkadink/ComfyUI-Advanced-ControlNet nodes, or make sure Advanced-ControlNet is updated.")
|
||||
prepare_control_objects(control_item, full_idxs)
|
||||
resized_actual_cond[key] = control_item
|
||||
del control_item
|
||||
elif isinstance(cond_item, dict):
|
||||
@@ -541,14 +564,13 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep,
|
||||
|
||||
if ADGS.motion_models is not None:
|
||||
ADGS.motion_models.set_view_options(ADGS.params.context_options.view_options)
|
||||
|
||||
# prepare final conds, out_counts, and biases
|
||||
conds_final = [torch.zeros_like(x_in) for _ in conds]
|
||||
counts_final = [torch.zeros((x_in.shape[0], 1, 1, 1), device=x_in.device) for _ in conds]
|
||||
biases_final = [([0.0] * x_in.shape[0]) for _ in conds]
|
||||
|
||||
# prepare final cond, uncond, and out_count
|
||||
cond_final = torch.zeros_like(x_in)
|
||||
uncond_final = torch.zeros_like(x_in)
|
||||
out_count_final = torch.zeros((x_in.shape[0], 1, 1, 1), device=x_in.device)
|
||||
bias_final = [0.0] * x_in.shape[0]
|
||||
|
||||
# perform calc_cond_uncond_batch per context window
|
||||
# perform calc_conds_batch per context window
|
||||
for ctx_idxs in context_windows:
|
||||
ADGS.params.sub_idxs = ctx_idxs
|
||||
if ADGS.motion_models is not None:
|
||||
@@ -562,16 +584,12 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep,
|
||||
for n in range(batched_conds):
|
||||
for ind in ctx_idxs:
|
||||
full_idxs.append((ADGS.params.full_length*n)+ind)
|
||||
# get subsections of x, timestep, cond, uncond, cond_concat
|
||||
# get subsections of x, timestep, conds
|
||||
sub_x = x_in[full_idxs]
|
||||
sub_timestep = timestep[full_idxs]
|
||||
sub_cond = get_resized_cond(cond, full_idxs, len(ctx_idxs)) if cond is not None else None
|
||||
sub_uncond = get_resized_cond(uncond, full_idxs, len(ctx_idxs)) if uncond is not None else None
|
||||
sub_conds = [get_resized_cond(cond, full_idxs, len(ctx_idxs)) for cond in conds]
|
||||
|
||||
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_conds_out = calc_cond_uncond_batch_wrapper(model, sub_conds, sub_x, sub_timestep, model_options)
|
||||
|
||||
if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE:
|
||||
full_length = ADGS.params.full_length
|
||||
@@ -581,28 +599,303 @@ def sliding_calc_cond_uncond_batch(model, cond, uncond, x_in: Tensor, timestep,
|
||||
bias = max(1e-2, bias)
|
||||
# take weighted average relative to total bias of current idx
|
||||
# and account for batched_conds
|
||||
for n in range(batched_conds):
|
||||
bias_total = bias_final[(full_length*n)+idx]
|
||||
prev_weight = (bias_total / (bias_total + bias))
|
||||
new_weight = (bias / (bias_total + bias))
|
||||
cond_final[(full_length*n)+idx] = cond_final[(full_length*n)+idx] * prev_weight + sub_cond_out[(full_length*n)+pos] * new_weight
|
||||
uncond_final[(full_length*n)+idx] = uncond_final[(full_length*n)+idx] * prev_weight + sub_uncond_out[(full_length*n)+pos] * new_weight
|
||||
bias_final[(full_length*n)+idx] = bias_total + bias
|
||||
for i in range(len(sub_conds_out)):
|
||||
for n in range(batched_conds):
|
||||
bias_total = biases_final[i][(full_length*n)+idx]
|
||||
prev_weight = (bias_total / (bias_total + bias))
|
||||
new_weight = (bias / (bias_total + bias))
|
||||
conds_final[i][(full_length*n)+idx] = conds_final[i][(full_length*n)+idx] * prev_weight + sub_conds_out[i][(full_length*n)+pos] * new_weight
|
||||
biases_final[i][(full_length*n)+idx] = bias_total + bias
|
||||
else:
|
||||
# add conds and counts based on weights of fuse method
|
||||
weights = get_context_weights(len(ctx_idxs), ADGS.params.context_options.fuse_method) * batched_conds
|
||||
weights_tensor = torch.Tensor(weights).to(device=x_in.device).unsqueeze(-1).unsqueeze(-1).unsqueeze(-1)
|
||||
cond_final[full_idxs] += sub_cond_out * weights_tensor
|
||||
uncond_final[full_idxs] += sub_uncond_out * weights_tensor
|
||||
out_count_final[full_idxs] += weights_tensor
|
||||
|
||||
for i in range(len(sub_conds_out)):
|
||||
conds_final[i][full_idxs] += sub_conds_out[i] * weights_tensor
|
||||
counts_final[i][full_idxs] += weights_tensor
|
||||
|
||||
if ADGS.params.context_options.fuse_method == ContextFuseMethod.RELATIVE:
|
||||
# already normalized, so return as is
|
||||
del out_count_final
|
||||
return cond_final, uncond_final
|
||||
del counts_final
|
||||
return conds_final
|
||||
else:
|
||||
# normalize cond and uncond via division by context usage counts
|
||||
cond_final /= out_count_final
|
||||
uncond_final /= out_count_final
|
||||
del out_count_final
|
||||
return cond_final, uncond_final
|
||||
# normalize conds via division by context usage counts
|
||||
for i in range(len(conds_final)):
|
||||
conds_final[i] /= counts_final[i]
|
||||
del counts_final
|
||||
return conds_final
|
||||
|
||||
|
||||
def calc_cond_uncond_batch_wrapper(model, conds: list[dict], x_in: Tensor, timestep, model_options):
|
||||
# check if conds or unconds contain lora_hook or default_cond
|
||||
contains_lora_hooks = False
|
||||
has_default_cond = False
|
||||
hook_groups = []
|
||||
for cond_uncond in conds:
|
||||
if cond_uncond is None:
|
||||
continue
|
||||
for t in cond_uncond:
|
||||
if COND_CONST.KEY_LORA_HOOK in t:
|
||||
contains_lora_hooks = True
|
||||
hook_groups.append(t[COND_CONST.KEY_LORA_HOOK])
|
||||
if COND_CONST.KEY_DEFAULT_COND in t:
|
||||
has_default_cond = True
|
||||
# if contains_lora_hooks:
|
||||
# break
|
||||
if contains_lora_hooks or has_default_cond:
|
||||
ADGS.hooks_initialize(model, hook_groups=hook_groups)
|
||||
ADGS.prepare_hooks_current_keyframes(timestep, hook_groups=hook_groups)
|
||||
return calc_conds_batch_lora_hook(model, conds, x_in, timestep, model_options, has_default_cond)
|
||||
# keep for backwards compatibility, for now
|
||||
if not hasattr(comfy.samplers, "calc_cond_batch"):
|
||||
return comfy.samplers.calc_cond_uncond_batch(model, conds[0], conds[1], x_in, timestep, model_options)
|
||||
return comfy.samplers.calc_cond_batch(model, conds, x_in, timestep, model_options)
|
||||
|
||||
|
||||
# modified from comfy.samplers.get_area_and_mult
|
||||
COND_OBJ = collections.namedtuple('cond_obj', ['input_x', 'mult', 'conditioning', 'area', 'control', 'patches'])
|
||||
def get_area_and_mult_ADE(conds, x_in, timestep_in):
|
||||
area = (x_in.shape[2], x_in.shape[3], 0, 0)
|
||||
strength = 1.0
|
||||
|
||||
if 'timestep_start' in conds:
|
||||
timestep_start = conds['timestep_start']
|
||||
if timestep_in[0] > timestep_start:
|
||||
return None
|
||||
if 'timestep_end' in conds:
|
||||
timestep_end = conds['timestep_end']
|
||||
if timestep_in[0] < timestep_end:
|
||||
return None
|
||||
if 'area' in conds:
|
||||
area = conds['area']
|
||||
if 'strength' in conds:
|
||||
strength = conds['strength']
|
||||
|
||||
input_x = x_in[:,:,area[2]:area[0] + area[2],area[3]:area[1] + area[3]]
|
||||
if 'mask' in conds:
|
||||
# Scale the mask to the size of the input
|
||||
# The mask should have been resized as we began the sampling process
|
||||
mask_strength = 1.0
|
||||
if "mask_strength" in conds:
|
||||
mask_strength = conds["mask_strength"]
|
||||
mask = conds['mask']
|
||||
assert(mask.shape[1] == x_in.shape[2])
|
||||
assert(mask.shape[2] == x_in.shape[3])
|
||||
# make sure mask is capped at input_shape batch length to prevent 0 as dimension
|
||||
mask = mask[:input_x.shape[0], area[2]:area[0] + area[2], area[3]:area[1] + area[3]] * mask_strength
|
||||
mask = mask.unsqueeze(1).repeat(input_x.shape[0] // mask.shape[0], input_x.shape[1], 1, 1)
|
||||
else:
|
||||
mask = torch.ones_like(input_x)
|
||||
mult = mask * strength
|
||||
|
||||
if 'mask' not in conds:
|
||||
rr = 8
|
||||
if area[2] != 0:
|
||||
for t in range(rr):
|
||||
mult[:,:,t:1+t,:] *= ((1.0/rr) * (t + 1))
|
||||
if (area[0] + area[2]) < x_in.shape[2]:
|
||||
for t in range(rr):
|
||||
mult[:,:,area[0] - 1 - t:area[0] - t,:] *= ((1.0/rr) * (t + 1))
|
||||
if area[3] != 0:
|
||||
for t in range(rr):
|
||||
mult[:,:,:,t:1+t] *= ((1.0/rr) * (t + 1))
|
||||
if (area[1] + area[3]) < x_in.shape[3]:
|
||||
for t in range(rr):
|
||||
mult[:,:,:,area[1] - 1 - t:area[1] - t] *= ((1.0/rr) * (t + 1))
|
||||
|
||||
conditioning = {}
|
||||
model_conds = conds["model_conds"]
|
||||
for c in model_conds:
|
||||
conditioning[c] = model_conds[c].process_cond(batch_size=x_in.shape[0], device=x_in.device, area=area)
|
||||
|
||||
control = conds.get('control', None)
|
||||
|
||||
patches = None
|
||||
if 'gligen' in conds:
|
||||
gligen = conds['gligen']
|
||||
patches = {}
|
||||
gligen_type = gligen[0]
|
||||
gligen_model = gligen[1]
|
||||
if gligen_type == "position":
|
||||
gligen_patch = gligen_model.model.set_position(input_x.shape, gligen[2], input_x.device)
|
||||
else:
|
||||
gligen_patch = gligen_model.model.set_empty(input_x.shape, input_x.device)
|
||||
|
||||
patches['middle_patch'] = [gligen_patch]
|
||||
|
||||
return COND_OBJ(input_x, mult, conditioning, area, control, patches)
|
||||
|
||||
|
||||
def separate_default_conds(conds: list[dict]):
|
||||
normal_conds = []
|
||||
default_conds = []
|
||||
for i in range(len(conds)):
|
||||
c = []
|
||||
default_c = []
|
||||
# if cond is None, make normal/default_conds reflect that too
|
||||
if conds[i] is None:
|
||||
c = None
|
||||
default_c = []
|
||||
else:
|
||||
for t in conds[i]:
|
||||
# check if cond is a default cond
|
||||
if COND_CONST.KEY_DEFAULT_COND in t:
|
||||
default_c.append(t)
|
||||
else:
|
||||
c.append(t)
|
||||
normal_conds.append(c)
|
||||
default_conds.append(default_c)
|
||||
return normal_conds, default_conds
|
||||
|
||||
|
||||
def finalize_default_conds(hooked_to_run: dict[LoraHookGroup,list[tuple[COND_OBJ,int]]], default_conds: list[list[dict]], x_in: Tensor, timestep):
|
||||
# need to figure out remaining unmasked area for conds
|
||||
default_mults = []
|
||||
for d in default_conds:
|
||||
default_mults.append(torch.ones_like(x_in))
|
||||
|
||||
# look through each finalized cond in hooked_to_run for 'mult' and subtract it from each cond
|
||||
for lora_hooks, to_run in hooked_to_run.items():
|
||||
for cond_obj, i in to_run:
|
||||
# if no default_cond for cond_type, do nothing
|
||||
if len(default_conds[i]) == 0:
|
||||
continue
|
||||
area: list[int] = cond_obj.area
|
||||
default_mults[i][:,:,area[2]:area[0] + area[2],area[3]:area[1] + area[3]] -= cond_obj.mult
|
||||
|
||||
# for each default_mult, ReLU to make negatives=0, and then check for any nonzeros
|
||||
for i, mult in enumerate(default_mults):
|
||||
# if no default_cond for cond type, do nothing
|
||||
if len(default_conds[i]) == 0:
|
||||
continue
|
||||
torch.nn.functional.relu(mult, inplace=True)
|
||||
# if mult is all zeros, then don't add default_cond
|
||||
if torch.max(mult) == 0.0:
|
||||
continue
|
||||
|
||||
cond = default_conds[i]
|
||||
for x in cond:
|
||||
# do get_area_and_mult to get all the expected values
|
||||
p = comfy.samplers.get_area_and_mult(x, x_in, timestep)
|
||||
if p is None:
|
||||
continue
|
||||
# replace p's mult with calculated mult
|
||||
p = p._replace(mult=mult)
|
||||
hook: LoraHookGroup = x.get(COND_CONST.KEY_LORA_HOOK, None)
|
||||
hooked_to_run.setdefault(hook, list())
|
||||
hooked_to_run[hook] += [(p, i)]
|
||||
|
||||
|
||||
# based on comfy.samplers.calc_conds_batch
|
||||
def calc_conds_batch_lora_hook(model: BaseModel, conds: list[list[dict]], x_in: Tensor, timestep, model_options: dict, has_default_cond=False):
|
||||
out_conds = []
|
||||
out_counts = []
|
||||
# separate conds by matching lora_hooks
|
||||
hooked_to_run: dict[LoraHookGroup,list[tuple[collections.namedtuple,int]]] = {}
|
||||
|
||||
# separate out default_conds, if needed
|
||||
if has_default_cond:
|
||||
conds, default_conds = separate_default_conds(conds)
|
||||
|
||||
# cond is i=0, uncond is i=1
|
||||
for i in range(len(conds)):
|
||||
out_conds.append(torch.zeros_like(x_in))
|
||||
out_counts.append(torch.ones_like(x_in) * 1e-37)
|
||||
|
||||
cond = conds[i]
|
||||
if cond is not None:
|
||||
for x in cond:
|
||||
p = comfy.samplers.get_area_and_mult(x, x_in, timestep)
|
||||
if p is None:
|
||||
continue
|
||||
hook: LoraHookGroup = x.get(COND_CONST.KEY_LORA_HOOK, None)
|
||||
hooked_to_run.setdefault(hook, list())
|
||||
hooked_to_run[hook] += [(p, i)]
|
||||
|
||||
# finalize default_conds, if needed
|
||||
if has_default_cond:
|
||||
finalize_default_conds(hooked_to_run, default_conds, x_in, timestep)
|
||||
|
||||
# run every hooked_to_run separately
|
||||
for lora_hooks, to_run in hooked_to_run.items():
|
||||
while len(to_run) > 0:
|
||||
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 = comfy.model_management.get_free_memory(x_in.device)
|
||||
for i in range(1, len(to_batch_temp) + 1):
|
||||
batch_amount = to_batch_temp[:len(to_batch_temp)//i]
|
||||
input_shape = [len(batch_amount) * first_shape[0]] + list(first_shape)[1:]
|
||||
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)
|
||||
|
||||
for o in range(batch_chunks):
|
||||
cond_index = cond_or_uncond[o]
|
||||
out_conds[cond_index][:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += output[o] * mult[o]
|
||||
out_counts[cond_index][:,:,area[o][2]:area[o][0] + area[o][2],area[o][3]:area[o][1] + area[o][3]] += mult[o]
|
||||
|
||||
for i in range(len(out_conds)):
|
||||
out_conds[i] /= out_counts[i]
|
||||
|
||||
return out_conds
|
||||
|
||||
@@ -157,7 +157,7 @@ def get_sorted_list_via_attr(objects: list, attr: str) -> list:
|
||||
unique_attrs = {}
|
||||
for o in objects:
|
||||
val_attr = getattr(o, attr)
|
||||
attr_list = unique_attrs.get(val_attr, list())
|
||||
attr_list: list = unique_attrs.get(val_attr, list())
|
||||
attr_list.append(o)
|
||||
if val_attr not in unique_attrs:
|
||||
unique_attrs[val_attr] = attr_list
|
||||
|
||||
Reference in New Issue
Block a user