diff --git a/README.md b/README.md index 802aa9c..a9eefe8 100644 --- a/README.md +++ b/README.md @@ -347,6 +347,8 @@ The parameters affect how the masked and unmasked prompts are combined to produc Creates a ComfyUI `HOOKS` object from a prompt schedule. Can be attached to a CLIP model to perform encoding and LoRA switching +The hooks can be a bit slow sometimes, especially in cases where the LoRA spec doesn't actually require reloading the patches every time you sample. Try `PCHooksFromScheduleWithOptimizationTest`. (Though that node is likely to go away eventually). + ## PCEncodeSchedule Encodes all prompts in a schedule. Pass in a `CLIP` object with hooks attached for LoRA scheduling, then use the resulting `CONDITIONING` normally diff --git a/__init__.py b/__init__.py index 59d057a..3a7c99e 100644 --- a/__init__.py +++ b/__init__.py @@ -30,10 +30,16 @@ else: import importlib if importlib.util.find_spec("comfy.hooks"): - from .prompt_control.node_hooks import PCLoraHooksFromSchedule, PCEncodeSchedule, PCEncodeSingle + from .prompt_control.node_hooks import ( + PCLoraHooksFromSchedule, + PCLoraHooksFromScheduleWithOptimizationTest, + PCEncodeSchedule, + PCEncodeSingle, + ) maps = { "PCLoraHooksFromSchedule": PCLoraHooksFromSchedule, + "PCLoraHooksFromScheduleWithOptimizationTest": PCLoraHooksFromScheduleWithOptimizationTest, "PCEncodeSchedule": PCEncodeSchedule, "PCEncodeSingle": PCEncodeSingle, } diff --git a/prompt_control/node_hooks.py b/prompt_control/node_hooks.py index 0a2ad27..5583e4f 100644 --- a/prompt_control/node_hooks.py +++ b/prompt_control/node_hooks.py @@ -4,6 +4,7 @@ import comfy.hooks import folder_paths from .prompts import encode_prompt from .utils import lora_name_to_file +import nodes log = logging.getLogger("comfyui-prompt-control") @@ -21,7 +22,64 @@ class PCLoraHooksFromSchedule: FUNCTION = "apply" def apply(self, prompt_schedule): - return (lora_hooks_from_schedule(prompt_schedule),) + consolidated = consolidate_schedule(prompt_schedule) + hooks = lora_hooks_from_schedule(consolidated, {}) + return (hooks,) + + +class PCLoraHooksFromScheduleWithOptimizationTest: + # Cache model + last_loras = None + last_modelclip = None + @classmethod + def INPUT_TYPES(s): + return { + "required": {"prompt_schedule": ("PROMPT_SCHEDULE",)}, + "optional": { + "model": ( + "MODEL", + { + "tooltip": "OPTIONAL model, passing it in enables an optimization to load the LoRA directly instead of creating a hook" + }, + ), + "clip": ("CLIP", {"tooltip": "OPTIONAL clip, like model"}), + }, + } + + RETURN_TYPES = ("HOOKS", "MODEL", "CLIP") + OUTPUT_TOOLTIPS = ( + "set of hooks created from the prompt schedule", + "optional output model with non-scheduled LoRAs applied, if an input model is provided", + "ditto for clip", + ) + CATEGORY = "promptcontrol/_unstable" + FUNCTION = "apply" + + def apply(self, prompt_schedule, model=None, clip=None): + + consolidated = consolidate_schedule(prompt_schedule) + + non_scheduled = {} + if model is not None: + loader = nodes.LoraLoader() + non_scheduled = find_nonscheduled_loras(consolidated) + if self.last_loras == non_scheduled and self.last_modelclip: + log.info("Returning cached models") + model, clip = self.last_modelclip + else: + for lora, info in non_scheduled.items(): + path = lora_name_to_file(lora) + if path is None: + log.info("LoRA not found: %s", lora) + continue + log.info("Attaching %s to model", lora) + model, clip = loader.load_lora(model, clip, path, info["weight"], info["weight_clip"]) + self.last_modelclip = model, clip + self.last_loras = non_scheduled + + hooks = lora_hooks_from_schedule(consolidated, non_scheduled) + + return hooks, model, clip class PCEncodeSchedule: @@ -55,17 +113,48 @@ class PCEncodeSingle: return (encode_prompt(clip, prompt, 0, 1.0, defaults or {}, None),) -def lora_hooks_from_schedule(schedules): +def consolidate_schedule(prompt_schedule): + prev_loras = {} + consolidated = [] + for end_pct, c in reversed(list(prompt_schedule)): + loras = c["loras"] + if loras != prev_loras: + consolidated.append((end_pct, loras)) + prev_loras = loras + return list(reversed(consolidated)) + + +def find_nonscheduled_loras(consolidated_schedule): + consolidated_schedule = list(consolidated_schedule) + if not consolidated_schedule: + return {} + last_end, candidate_loras = consolidated_schedule[0] + print(candidate_loras) + to_remove = set() + for candidate, weights in candidate_loras.items(): + for end, loras in consolidated_schedule[1:]: + last_end = end + if loras.get(candidate) != weights: + to_remove.add(candidate) + # No candidates if the schedule does not span full time + if last_end < 1.0: + return {} + return {k: v for (k, v) in candidate_loras.items() if k not in to_remove} + + +def lora_hooks_from_schedule(schedules, non_scheduled): start_pct = 0.0 lora_cache = {} all_hooks = [] - prev_loras = {} - def create_hook(loraspec, start_pct, end_pct): + def create_hook(loraspec, start_pct, end_pct, non_scheduled): nonlocal lora_cache hooks = [] hook_kf = comfy.hooks.HookKeyframeGroup() for lora, info in loras.items(): + if non_scheduled.get(lora) == info: + log.info("Skipping %s from hook, it's loaded directly on model", lora) + continue path = lora_name_to_file(lora) if not path: continue @@ -92,25 +181,15 @@ def lora_hooks_from_schedule(schedules): hooks.set_keyframes_on_hooks(hook_kf=hook_kf) return hooks - consolidated = [] - - prev_loras = {} - for end_pct, c in reversed(list(schedules)): - loras = c["loras"] - if loras != prev_loras: - consolidated.append((end_pct, loras)) - prev_loras = loras - consolidated = reversed(consolidated) - - for end_pct, loras in consolidated: + for end_pct, loras in schedules: log.info("Creating LoRA hook from %s to %s: %s", start_pct, end_pct, loras) - hook = create_hook(loras, start_pct, end_pct) + hook = create_hook(loras, start_pct, end_pct, non_scheduled) all_hooks.append(hook) start_pct = end_pct del lora_cache - all_hooks = [x for x in all_hooks if x is not None] + all_hooks = [x for x in all_hooks if x] if all_hooks: hooks = comfy.hooks.HookGroup.combine_all_hooks(all_hooks)