Add LoRA loading optimization. I don't think I like this interface though...

This commit is contained in:
asagi4
2024-12-07 21:11:38 +02:00
parent dab719f369
commit 58fe45eb87
3 changed files with 105 additions and 18 deletions
+2
View File
@@ -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
View File
@@ -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,
}
+96 -17
View File
@@ -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)