From d5169df808b00de23d651a5e63d2c00b57257d66 Mon Sep 17 00:00:00 2001 From: Jedrzej Kosinski Date: Wed, 30 Oct 2024 04:56:09 -0500 Subject: [PATCH] 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 --- comfy/hooks.py | 54 +++++++++++++++++++++++++++++++++++++ comfy/model_patcher.py | 2 +- comfy/samplers.py | 14 ++++++---- comfy/sd.py | 39 +++++++++++++++++++++++++++ comfy_extras/nodes_hooks.py | 11 +++++--- nodes.py | 8 +++--- 6 files changed, 115 insertions(+), 13 deletions(-) diff --git a/comfy/hooks.py b/comfy/hooks.py index e5fd07599..85dcd9281 100644 --- a/comfy/hooks.py +++ b/comfy/hooks.py @@ -4,6 +4,7 @@ import enum import math import torch import numpy as np +import itertools if TYPE_CHECKING: from comfy.model_patcher import ModelPatcher, PatcherInjection @@ -278,6 +279,59 @@ class HookGroup: with_type[hook] = None 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 def combine_all_hooks(hooks_list: list['HookGroup'], require_count=0) -> 'HookGroup': actual: list[HookGroup] = [] diff --git a/comfy/model_patcher.py b/comfy/model_patcher.py index 00ac4d14b..ef7971b14 100644 --- a/comfy/model_patcher.py +++ b/comfy/model_patcher.py @@ -971,7 +971,7 @@ class ModelPatcher: if hook.hook_ref not in self.hook_patches: weight_hooks_to_register.append(hook) 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: hook.add_hook_patches(self, target, registered_hooks) for callback in self.get_all_callbacks(CallbacksMP.ON_REGISTER_ALL_HOOK_PATCHES): diff --git a/comfy/samplers.py b/comfy/samplers.py index cbcd350c5..1a9cd1e58 100644 --- a/comfy/samplers.py +++ b/comfy/samplers.py @@ -195,7 +195,6 @@ def outer_calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: tor out_conds = [] out_counts = [] # separate conds by matching hooks - # TODO: implement default_conds support hooked_to_run: dict[comfy.hooks.HookGroup,list[tuple[tuple,int]]] = {} default_conds = [] has_default_conds = False @@ -557,10 +556,15 @@ def calculate_start_end_timesteps(model, conds): timestep_start = None timestep_end = None - 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']) + # handle clip hook schedule, if needed + if 'clip_start_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(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): n = x.copy() diff --git a/comfy/sd.py b/comfy/sd.py index e4abf0b94..42f805d3b 100644 --- a/comfy/sd.py +++ b/comfy/sd.py @@ -1,3 +1,4 @@ +from __future__ import annotations import torch from enum import Enum import logging @@ -28,6 +29,7 @@ import comfy.text_encoders.long_clipl import comfy.model_patcher import comfy.lora +import comfy.hooks import comfy.t2i_adapter.adapter import comfy.taesd.taesd @@ -90,9 +92,11 @@ class CLIP: 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.hook_mode = comfy.hooks.EnumHookMode.MinVram if params['device'] == load_device: model_management.load_models_gpu([self.patcher], force_full_load=True) 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'])) def clone(self): @@ -101,6 +105,7 @@ class CLIP: n.cond_stage_model = self.cond_stage_model n.tokenizer = self.tokenizer n.layer_idx = self.layer_idx + n.use_clip_schedule = self.use_clip_schedule return n 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): 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): self.cond_stage_model.reset_clip_options() diff --git a/comfy_extras/nodes_hooks.py b/comfy_extras/nodes_hooks.py index 420259785..8d933ba08 100644 --- a/comfy_extras/nodes_hooks.py +++ b/comfy_extras/nodes_hooks.py @@ -228,9 +228,10 @@ class SetClipHooks: return { "required": { "clip": ("CLIP",), + "schedule_clip": ("BOOLEAN", {"default": True}) }, "optional": { - "hooks": ("HOOKS",), + "hooks": ("HOOKS",) } } @@ -238,9 +239,10 @@ class SetClipHooks: CATEGORY = "advanced/hooks/clip" 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: clip = clip.clone() + clip.use_clip_schedule = schedule_clip clip.patcher.forced_hooks = hooks clip.patcher.register_all_hook_patches(hooks.get_dict_repr(), comfy.hooks.EnumWeightTarget.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" FUNCTION = "create_range" 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)) #------------------------------------------ ########################################### diff --git a/nodes.py b/nodes.py index 0e38489a6..498892517 100644 --- a/nodes.py +++ b/nodes.py @@ -62,9 +62,11 @@ class CLIPTextEncode: def encode(self, clip, text): tokens = clip.tokenize(text) - output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True) - cond = output.pop("cond") - return ([[cond, output]], ) + if not clip.use_clip_schedule: + output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True) + cond = output.pop("cond") + return ([[cond, output]], ) + return (clip.encode_from_tokens_scheduled(tokens), ) class ConditioningCombine: @classmethod