Add LoRA loading optimization. I don't think I like this interface though...
This commit is contained in:
@@ -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
|
||||
|
||||
+7
-1
@@ -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,
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user