132 lines
4.7 KiB
Python
132 lines
4.7 KiB
Python
import logging
|
|
|
|
import comfy.hooks
|
|
import comfy.utils
|
|
import folder_paths
|
|
from comfy_api.latest import io
|
|
from typing_extensions import override
|
|
|
|
from .attention_couple_ppm import AttentionCoupleHook
|
|
from .parser import parse_prompt_schedules
|
|
from .utils import consolidate_schedule
|
|
|
|
log = logging.getLogger("comfyui-prompt-control")
|
|
|
|
|
|
class PCLoraHooksFromText(io.ComfyNode):
|
|
@classmethod
|
|
def define_schema(cls):
|
|
return io.Schema(
|
|
node_id="PCLoraHooksFromText",
|
|
display_name="PC: LoRA Hooks From Text (non-lazy)",
|
|
category="promptcontrol/v2",
|
|
description="set of hooks created from the prompt schedule",
|
|
is_experimental=True,
|
|
inputs=[
|
|
io.String.Input("text", multiline=True),
|
|
],
|
|
outputs=[io.Hooks.Output()],
|
|
)
|
|
|
|
@classmethod
|
|
def execute(cls, text) -> io.NodeOutput:
|
|
prompt_schedule = parse_prompt_schedules(text)
|
|
consolidated = consolidate_schedule(prompt_schedule)
|
|
hooks = lora_hooks_from_schedule(consolidated, {})
|
|
return io.NodeOutput(hooks)
|
|
|
|
|
|
def lora_hooks_from_schedule(schedules, non_scheduled):
|
|
start_pct = 0.0
|
|
lora_cache = {}
|
|
all_hooks = []
|
|
|
|
def create_hook(loras, start_pct, end_pct, non_scheduled):
|
|
hooks = []
|
|
hook_kf = comfy.hooks.HookKeyframeGroup()
|
|
for path, info in loras.items():
|
|
if non_scheduled.get(path) == info:
|
|
log.info("Skipping %s from hook, it's loaded directly on model", path)
|
|
continue
|
|
if path not in lora_cache:
|
|
lora_cache[path] = comfy.utils.load_torch_file(
|
|
folder_paths.get_full_path("loras", path), safe_load=True
|
|
)
|
|
new_hook = comfy.hooks.create_hook_lora(
|
|
lora_cache[path], strength_model=info["weight"], strength_clip=info["weight_clip"]
|
|
)
|
|
ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
|
new_hook.hooks[0].hook_ref = ref
|
|
hooks.append(new_hook)
|
|
if start_pct > 0.0:
|
|
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
|
|
hook_kf.add(kf)
|
|
kf = comfy.hooks.HookKeyframe(strength=1.0, start_percent=start_pct)
|
|
hook_kf.add(kf)
|
|
if end_pct < 1.0:
|
|
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=end_pct)
|
|
hook_kf.add(kf)
|
|
hooks = comfy.hooks.HookGroup.combine_all_hooks(hooks)
|
|
if hooks:
|
|
hooks.set_keyframes_on_hooks(hook_kf=hook_kf)
|
|
return hooks
|
|
|
|
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, non_scheduled)
|
|
all_hooks.append(hook)
|
|
start_pct = end_pct
|
|
|
|
all_hooks = [x for x in all_hooks if x]
|
|
if all_hooks:
|
|
hooks = comfy.hooks.HookGroup.combine_all_hooks(all_hooks)
|
|
return hooks
|
|
|
|
|
|
class PCAttentionCoupleBatchNegative(io.ComfyNode):
|
|
@classmethod
|
|
def define_schema(cls):
|
|
return io.Schema(
|
|
node_id="PCAttentionCoupleBatchNegative",
|
|
display_name="PC: Attention Couple (batch negative)",
|
|
category="promptcontrol/v2",
|
|
description="Batch negatives, carrying over Attention Couple hooks",
|
|
is_experimental=True,
|
|
inputs=[
|
|
io.Conditioning.Input("positive"),
|
|
io.Conditioning.Input("negative"),
|
|
],
|
|
outputs=[
|
|
io.Conditioning.Output("positive"),
|
|
io.Conditioning.Output("negative"),
|
|
],
|
|
)
|
|
|
|
@classmethod
|
|
@override
|
|
def execute(cls, positive, negative) -> io.NodeOutput:
|
|
if len(negative) != 1:
|
|
log.warning("Batching scheduled negatives is not supported yet")
|
|
return io.NodeOutput(positive, negative)
|
|
|
|
negative_batch = []
|
|
for p in positive:
|
|
n = [negative[0][0], negative[0][1].copy()]
|
|
n_hook_group = n[1].get("hooks", comfy.hooks.HookGroup()).clone()
|
|
p_hook_group = p[1].get("hooks", comfy.hooks.HookGroup())
|
|
attn_couple = [hook for hook in p_hook_group.hooks if isinstance(hook, AttentionCoupleHook)]
|
|
for hook in attn_couple:
|
|
n_hook_group.add(hook)
|
|
n[1]["hooks"] = p_hook_group if n_hook_group.hooks == p_hook_group.hooks else n_hook_group
|
|
n[1]["start_percent"] = p[1].get("start_percent", 0.0)
|
|
n[1]["end_percent"] = p[1].get("end_percent", 1.0)
|
|
negative_batch.append(n)
|
|
|
|
return io.NodeOutput(positive, negative_batch)
|
|
|
|
|
|
NODES = [
|
|
PCLoraHooksFromText,
|
|
PCAttentionCoupleBatchNegative,
|
|
]
|