mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-10-08 14:57:04 +08:00
Added initial support within CLIP Text Encode (Prompt) node for scheduling weight hook CLIP strength via clip_start_percent/clip_end_percent on conds, added schedule_clip toggle to Set CLIP Hooks node, small cleanup/fixes
This commit is contained in:
parent
2047bf211f
commit
d5169df808
@ -4,6 +4,7 @@ import enum
|
|||||||
import math
|
import math
|
||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import itertools
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from comfy.model_patcher import ModelPatcher, PatcherInjection
|
from comfy.model_patcher import ModelPatcher, PatcherInjection
|
||||||
@ -278,6 +279,59 @@ class HookGroup:
|
|||||||
with_type[hook] = None
|
with_type[hook] = None
|
||||||
return d
|
return d
|
||||||
|
|
||||||
|
def get_hooks_for_clip_schedule(self):
|
||||||
|
scheduled_hooks: dict[WeightHook, list[tuple[tuple[float,float], HookKeyframe]]] = {}
|
||||||
|
for hook in self.hooks:
|
||||||
|
# only care about WeightHooks, for now
|
||||||
|
if hook.hook_type == EnumHookType.Weight:
|
||||||
|
hook_schedule = []
|
||||||
|
# if no hook keyframes, assign default value
|
||||||
|
if len(hook.hook_keyframe.keyframes) == 0:
|
||||||
|
hook_schedule.append(((0.0, 1.0), None))
|
||||||
|
scheduled_hooks[hook] = hook_schedule
|
||||||
|
continue
|
||||||
|
# find ranges of values
|
||||||
|
prev_keyframe = hook.hook_keyframe.keyframes[0]
|
||||||
|
for keyframe in hook.hook_keyframe.keyframes:
|
||||||
|
if keyframe.start_percent > prev_keyframe.start_percent:
|
||||||
|
hook_schedule.append(((prev_keyframe.start_percent, keyframe.start_percent), prev_keyframe))
|
||||||
|
prev_keyframe = keyframe
|
||||||
|
elif keyframe.start_percent == prev_keyframe.start_percent:
|
||||||
|
prev_keyframe = keyframe
|
||||||
|
# create final range, assuming last start_percent was not 1.0
|
||||||
|
if not math.isclose(prev_keyframe.start_percent, 1.0):
|
||||||
|
hook_schedule.append(((prev_keyframe.start_percent, 1.0), prev_keyframe))
|
||||||
|
scheduled_hooks[hook] = hook_schedule
|
||||||
|
# hooks should not have their schedules in a list of tuples
|
||||||
|
all_ranges: list[tuple[float, float]] = []
|
||||||
|
for range_kfs in scheduled_hooks.values():
|
||||||
|
for t_range, keyframe in range_kfs:
|
||||||
|
all_ranges.append(t_range)
|
||||||
|
# turn list of ranges into boundaries
|
||||||
|
boundaries_set = set(itertools.chain.from_iterable(all_ranges))
|
||||||
|
boundaries_set.add(0.0)
|
||||||
|
boundaries = sorted(boundaries_set)
|
||||||
|
real_ranges = [(boundaries[i], boundaries[i + 1]) for i in range(len(boundaries) - 1)]
|
||||||
|
# with real ranges defined, give appropriate hooks w/ keyframes for each range
|
||||||
|
scheduled_keyframes: list[tuple[tuple[float,float], list[tuple[WeightHook, HookKeyframe]]]] = []
|
||||||
|
for t_range in real_ranges:
|
||||||
|
hooks_schedule = []
|
||||||
|
for hook, val in scheduled_hooks.items():
|
||||||
|
keyframe = None
|
||||||
|
# check if is a keyframe that works for the current t_range
|
||||||
|
for stored_range, stored_kf in val:
|
||||||
|
# if stored start is less than current end, then fits - give it assigned keyframe
|
||||||
|
if stored_range[0] < t_range[1]:
|
||||||
|
keyframe = stored_kf
|
||||||
|
break
|
||||||
|
hooks_schedule.append((hook, keyframe))
|
||||||
|
scheduled_keyframes.append((t_range, hooks_schedule))
|
||||||
|
return scheduled_keyframes
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
for hook in self.hooks:
|
||||||
|
hook.hook_keyframe.reset()
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def combine_all_hooks(hooks_list: list['HookGroup'], require_count=0) -> 'HookGroup':
|
def combine_all_hooks(hooks_list: list['HookGroup'], require_count=0) -> 'HookGroup':
|
||||||
actual: list[HookGroup] = []
|
actual: list[HookGroup] = []
|
||||||
|
|||||||
@ -971,7 +971,7 @@ class ModelPatcher:
|
|||||||
if hook.hook_ref not in self.hook_patches:
|
if hook.hook_ref not in self.hook_patches:
|
||||||
weight_hooks_to_register.append(hook)
|
weight_hooks_to_register.append(hook)
|
||||||
if len(weight_hooks_to_register) > 0:
|
if len(weight_hooks_to_register) > 0:
|
||||||
self.hook_patches_backup = create_hook_patches_clone(self.hook_patches)
|
self.hook_patches_backup = create_hook_patches_clone(self.hook_patches_backup)
|
||||||
for hook in weight_hooks_to_register:
|
for hook in weight_hooks_to_register:
|
||||||
hook.add_hook_patches(self, target, registered_hooks)
|
hook.add_hook_patches(self, target, registered_hooks)
|
||||||
for callback in self.get_all_callbacks(CallbacksMP.ON_REGISTER_ALL_HOOK_PATCHES):
|
for callback in self.get_all_callbacks(CallbacksMP.ON_REGISTER_ALL_HOOK_PATCHES):
|
||||||
|
|||||||
@ -195,7 +195,6 @@ def outer_calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: tor
|
|||||||
out_conds = []
|
out_conds = []
|
||||||
out_counts = []
|
out_counts = []
|
||||||
# separate conds by matching hooks
|
# separate conds by matching hooks
|
||||||
# TODO: implement default_conds support
|
|
||||||
hooked_to_run: dict[comfy.hooks.HookGroup,list[tuple[tuple,int]]] = {}
|
hooked_to_run: dict[comfy.hooks.HookGroup,list[tuple[tuple,int]]] = {}
|
||||||
default_conds = []
|
default_conds = []
|
||||||
has_default_conds = False
|
has_default_conds = False
|
||||||
@ -557,10 +556,15 @@ def calculate_start_end_timesteps(model, conds):
|
|||||||
|
|
||||||
timestep_start = None
|
timestep_start = None
|
||||||
timestep_end = None
|
timestep_end = None
|
||||||
if 'start_percent' in x:
|
# handle clip hook schedule, if needed
|
||||||
timestep_start = s.percent_to_sigma(x['start_percent'])
|
if 'clip_start_percent' in x:
|
||||||
if 'end_percent' in x:
|
timestep_start = s.percent_to_sigma(max(x['clip_start_percent'], x.get('start_percent', 0.0)))
|
||||||
timestep_end = s.percent_to_sigma(x['end_percent'])
|
timestep_end = s.percent_to_sigma(min(x['clip_end_percent'], x.get('end_percent', 1.0)))
|
||||||
|
else:
|
||||||
|
if 'start_percent' in x:
|
||||||
|
timestep_start = s.percent_to_sigma(x['start_percent'])
|
||||||
|
if 'end_percent' in x:
|
||||||
|
timestep_end = s.percent_to_sigma(x['end_percent'])
|
||||||
|
|
||||||
if (timestep_start is not None) or (timestep_end is not None):
|
if (timestep_start is not None) or (timestep_end is not None):
|
||||||
n = x.copy()
|
n = x.copy()
|
||||||
|
|||||||
39
comfy/sd.py
39
comfy/sd.py
@ -1,3 +1,4 @@
|
|||||||
|
from __future__ import annotations
|
||||||
import torch
|
import torch
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
import logging
|
import logging
|
||||||
@ -28,6 +29,7 @@ import comfy.text_encoders.long_clipl
|
|||||||
|
|
||||||
import comfy.model_patcher
|
import comfy.model_patcher
|
||||||
import comfy.lora
|
import comfy.lora
|
||||||
|
import comfy.hooks
|
||||||
import comfy.t2i_adapter.adapter
|
import comfy.t2i_adapter.adapter
|
||||||
import comfy.taesd.taesd
|
import comfy.taesd.taesd
|
||||||
|
|
||||||
@ -90,9 +92,11 @@ class CLIP:
|
|||||||
|
|
||||||
self.tokenizer = tokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data)
|
self.tokenizer = tokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data)
|
||||||
self.patcher = comfy.model_patcher.ModelPatcher(self.cond_stage_model, load_device=load_device, offload_device=offload_device)
|
self.patcher = comfy.model_patcher.ModelPatcher(self.cond_stage_model, load_device=load_device, offload_device=offload_device)
|
||||||
|
self.patcher.hook_mode = comfy.hooks.EnumHookMode.MinVram
|
||||||
if params['device'] == load_device:
|
if params['device'] == load_device:
|
||||||
model_management.load_models_gpu([self.patcher], force_full_load=True)
|
model_management.load_models_gpu([self.patcher], force_full_load=True)
|
||||||
self.layer_idx = None
|
self.layer_idx = None
|
||||||
|
self.use_clip_schedule = False
|
||||||
logging.debug("CLIP model load device: {}, offload device: {}, current: {}".format(load_device, offload_device, params['device']))
|
logging.debug("CLIP model load device: {}, offload device: {}, current: {}".format(load_device, offload_device, params['device']))
|
||||||
|
|
||||||
def clone(self):
|
def clone(self):
|
||||||
@ -101,6 +105,7 @@ class CLIP:
|
|||||||
n.cond_stage_model = self.cond_stage_model
|
n.cond_stage_model = self.cond_stage_model
|
||||||
n.tokenizer = self.tokenizer
|
n.tokenizer = self.tokenizer
|
||||||
n.layer_idx = self.layer_idx
|
n.layer_idx = self.layer_idx
|
||||||
|
n.use_clip_schedule = self.use_clip_schedule
|
||||||
return n
|
return n
|
||||||
|
|
||||||
def add_patches(self, patches, strength_patch=1.0, strength_model=1.0):
|
def add_patches(self, patches, strength_patch=1.0, strength_model=1.0):
|
||||||
@ -112,6 +117,40 @@ class CLIP:
|
|||||||
def tokenize(self, text, return_word_ids=False):
|
def tokenize(self, text, return_word_ids=False):
|
||||||
return self.tokenizer.tokenize_with_weights(text, return_word_ids)
|
return self.tokenizer.tokenize_with_weights(text, return_word_ids)
|
||||||
|
|
||||||
|
def encode_from_tokens_scheduled(self, tokens, unprojected=False, add_dict: dict[str]=None):
|
||||||
|
all_cond_pooled: list[tuple[torch.Tensor, dict[str]]] = []
|
||||||
|
all_hooks = self.patcher.forced_hooks
|
||||||
|
scheduled_keyframes = all_hooks.get_hooks_for_clip_schedule()
|
||||||
|
|
||||||
|
self.cond_stage_model.reset_clip_options()
|
||||||
|
if self.layer_idx is not None:
|
||||||
|
self.cond_stage_model.set_clip_options({"layer": self.layer_idx})
|
||||||
|
if unprojected:
|
||||||
|
self.cond_stage_model.set_clip_options({"projected_pooled": False})
|
||||||
|
|
||||||
|
self.load_model()
|
||||||
|
all_hooks.reset()
|
||||||
|
for scheduled_opts in scheduled_keyframes:
|
||||||
|
t_range = scheduled_opts[0]
|
||||||
|
hooks_keyframes = scheduled_opts[1]
|
||||||
|
for hook, keyframe in hooks_keyframes:
|
||||||
|
hook.hook_keyframe._current_keyframe = keyframe
|
||||||
|
# apply appropriate hooks with values that match new hook_keyframe
|
||||||
|
self.patcher.patch_hooks(all_hooks)
|
||||||
|
# perform encoding as normal
|
||||||
|
o = self.cond_stage_model.encode_token_weights(tokens)
|
||||||
|
cond, pooled = o[:2]
|
||||||
|
pooled_dict = {"pooled_output": pooled}
|
||||||
|
# add clip_start_percent and clip_end_percent in pooled
|
||||||
|
pooled_dict["clip_start_percent"] = t_range[0]
|
||||||
|
pooled_dict["clip_end_percent"] = t_range[1]
|
||||||
|
# add/update any keys with the provided add_dict
|
||||||
|
if add_dict is not None:
|
||||||
|
pooled_dict.update(add_dict)
|
||||||
|
all_cond_pooled.append([cond, pooled_dict])
|
||||||
|
all_hooks.reset()
|
||||||
|
return all_cond_pooled
|
||||||
|
|
||||||
def encode_from_tokens(self, tokens, return_pooled=False, return_dict=False):
|
def encode_from_tokens(self, tokens, return_pooled=False, return_dict=False):
|
||||||
self.cond_stage_model.reset_clip_options()
|
self.cond_stage_model.reset_clip_options()
|
||||||
|
|
||||||
|
|||||||
@ -228,9 +228,10 @@ class SetClipHooks:
|
|||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"clip": ("CLIP",),
|
"clip": ("CLIP",),
|
||||||
|
"schedule_clip": ("BOOLEAN", {"default": True})
|
||||||
},
|
},
|
||||||
"optional": {
|
"optional": {
|
||||||
"hooks": ("HOOKS",),
|
"hooks": ("HOOKS",)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -238,9 +239,10 @@ class SetClipHooks:
|
|||||||
CATEGORY = "advanced/hooks/clip"
|
CATEGORY = "advanced/hooks/clip"
|
||||||
FUNCTION = "apply_hooks"
|
FUNCTION = "apply_hooks"
|
||||||
|
|
||||||
def apply_hooks(self, clip: 'CLIP', hooks: comfy.hooks.HookGroup=None):
|
def apply_hooks(self, clip: 'CLIP', schedule_clip: bool, hooks: comfy.hooks.HookGroup=None):
|
||||||
if hooks is not None:
|
if hooks is not None:
|
||||||
clip = clip.clone()
|
clip = clip.clone()
|
||||||
|
clip.use_clip_schedule = schedule_clip
|
||||||
clip.patcher.forced_hooks = hooks
|
clip.patcher.forced_hooks = hooks
|
||||||
clip.patcher.register_all_hook_patches(hooks.get_dict_repr(), comfy.hooks.EnumWeightTarget.Clip)
|
clip.patcher.register_all_hook_patches(hooks.get_dict_repr(), comfy.hooks.EnumWeightTarget.Clip)
|
||||||
return (clip,)
|
return (clip,)
|
||||||
@ -257,12 +259,13 @@ class ConditioningTimestepsRange:
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("TIMESTEPS_RANGE",)
|
RETURN_TYPES = ("TIMESTEPS_RANGE", "TIMESTEPS_RANGE", "TIMESTEPS_RANGE")
|
||||||
|
RETURN_NAMES = ("TIMESTEPS_RANGE", "BEFORE_RANGE", "AFTER_RANGE")
|
||||||
CATEGORY = "advanced/hooks"
|
CATEGORY = "advanced/hooks"
|
||||||
FUNCTION = "create_range"
|
FUNCTION = "create_range"
|
||||||
|
|
||||||
def create_range(self, start_percent: float, end_percent: float):
|
def create_range(self, start_percent: float, end_percent: float):
|
||||||
return ((start_percent, end_percent),)
|
return ((start_percent, end_percent), (0.0, start_percent), (end_percent, 1.0))
|
||||||
#------------------------------------------
|
#------------------------------------------
|
||||||
###########################################
|
###########################################
|
||||||
|
|
||||||
|
|||||||
8
nodes.py
8
nodes.py
@ -62,9 +62,11 @@ class CLIPTextEncode:
|
|||||||
|
|
||||||
def encode(self, clip, text):
|
def encode(self, clip, text):
|
||||||
tokens = clip.tokenize(text)
|
tokens = clip.tokenize(text)
|
||||||
output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True)
|
if not clip.use_clip_schedule:
|
||||||
cond = output.pop("cond")
|
output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True)
|
||||||
return ([[cond, output]], )
|
cond = output.pop("cond")
|
||||||
|
return ([[cond, output]], )
|
||||||
|
return (clip.encode_from_tokens_scheduled(tokens), )
|
||||||
|
|
||||||
class ConditioningCombine:
|
class ConditioningCombine:
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user