Clean up legacy stuff
This commit is contained in:
@@ -1,8 +1,8 @@
|
||||
all: format check
|
||||
@echo "Done"
|
||||
check:
|
||||
pyflakes *.py */*.py */*/*.py
|
||||
pyflakes *.py */*.py
|
||||
format:
|
||||
black -l 120 *.py */*.py */*/*.py
|
||||
black -l 120 *.py */*.py
|
||||
|
||||
.PHONY: check format all
|
||||
|
||||
+7
-33
@@ -2,20 +2,6 @@ import os
|
||||
import sys
|
||||
import logging
|
||||
|
||||
from .prompt_control.legacy.node_clip import EditableCLIPEncode, ScheduleToCond
|
||||
from .prompt_control.legacy.node_lora import LoRAScheduler, ScheduleToModel, PCSplitSampling, PCWrapGuider
|
||||
from .prompt_control.legacy.node_other import (
|
||||
PromptToSchedule,
|
||||
FilterSchedule,
|
||||
PCScheduleSettings,
|
||||
PCScheduleAddMasks,
|
||||
PCApplySettings,
|
||||
PCPromptFromSchedule,
|
||||
)
|
||||
from .prompt_control.legacy.node_aio import PromptControlSimple
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
log.propagate = False
|
||||
@@ -29,9 +15,16 @@ if os.environ.get("COMFYUI_PC_DEBUG"):
|
||||
else:
|
||||
log.setLevel(logging.INFO)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
from .prompt_control.node_other import NODE_CLASS_MAPPINGS as o_mappings, NODE_DISPLAY_NAME_MAPPINGS as o_display
|
||||
from .prompt_control.nodes_lazy import NODE_CLASS_MAPPINGS as lazy_mappings, NODE_DISPLAY_NAME_MAPPINGS as lazy_display
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(o_mappings)
|
||||
NODE_CLASS_MAPPINGS.update(lazy_mappings)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(o_display)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(lazy_display)
|
||||
|
||||
import importlib
|
||||
@@ -50,22 +43,3 @@ else:
|
||||
)
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.realpath(__file__)), "comfy"))
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(
|
||||
{
|
||||
"PromptControlSimple": PromptControlSimple,
|
||||
"PromptToSchedule": PromptToSchedule,
|
||||
"PCSplitSampling": PCSplitSampling,
|
||||
"PCScheduleSettings": PCScheduleSettings,
|
||||
"PCScheduleAddMasks": PCScheduleAddMasks,
|
||||
"PCApplySettings": PCApplySettings,
|
||||
"PCPromptFromSchedule": PCPromptFromSchedule,
|
||||
"PCWrapGuider": PCWrapGuider,
|
||||
"FilterSchedule": FilterSchedule,
|
||||
"ScheduleToCond": ScheduleToCond,
|
||||
"ScheduleToModel": ScheduleToModel,
|
||||
"EditableCLIPEncode": EditableCLIPEncode,
|
||||
"LoRAScheduler": LoRAScheduler,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -1,160 +0,0 @@
|
||||
from .utils import get_callback, unpatch_model
|
||||
import sys
|
||||
|
||||
import logging
|
||||
import gc
|
||||
import comfy.model_management
|
||||
import os
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def has_hijack(obj):
|
||||
return hasattr(obj, "pc_hijack_done")
|
||||
|
||||
|
||||
def hijack(obj, attr, replacement):
|
||||
setattr(obj, attr, replacement)
|
||||
setattr(replacement, "pc_hijack_done", True)
|
||||
|
||||
|
||||
def hijack_sampler(module, function, is_custom):
|
||||
mod = sys.modules[module]
|
||||
orig_sampler = getattr(mod, function)
|
||||
if has_hijack(orig_sampler):
|
||||
return
|
||||
|
||||
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler
|
||||
|
||||
def pc_sample(*args, **kwargs):
|
||||
model = args[0]
|
||||
cb = get_callback(model)
|
||||
BrownianTreeNoiseSampler.pc_reset(
|
||||
model.model_options.get("pc_split_sampling"),
|
||||
kwargs.get("force_full_denoise") or kwargs.get("denoise", 1.0) >= 1.0,
|
||||
)
|
||||
if cb:
|
||||
try:
|
||||
try:
|
||||
r = cb(orig_sampler, is_custom, *args, **kwargs)
|
||||
except comfy.model_management.OOM_EXCEPTION:
|
||||
if not os.environ.get("PC_RETRY_ON_OOM"):
|
||||
raise
|
||||
log.error("Got OOM while sampling, freeing memory and retrying once...")
|
||||
unpatch_model(model)
|
||||
BrownianTreeNoiseSampler.pc_reset(False)
|
||||
gc.collect()
|
||||
comfy.model_management.soft_empty_cache()
|
||||
r = cb(orig_sampler, is_custom, *args, **kwargs)
|
||||
except Exception:
|
||||
log.error("Exception occurred during callback, unpatching model.")
|
||||
unpatch_model(model)
|
||||
BrownianTreeNoiseSampler.pc_reset(False)
|
||||
raise
|
||||
else:
|
||||
r = orig_sampler(*args, **kwargs)
|
||||
BrownianTreeNoiseSampler.pc_reset()
|
||||
return r
|
||||
|
||||
hijack(mod, function, pc_sample)
|
||||
|
||||
|
||||
def hijack_ksampler(module, cls):
|
||||
mod = sys.modules[module]
|
||||
orig_sampler = getattr(mod, cls)
|
||||
if has_hijack(orig_sampler):
|
||||
return
|
||||
|
||||
class HijackedKSampler(orig_sampler):
|
||||
def sample(
|
||||
self,
|
||||
noise,
|
||||
positive,
|
||||
negative,
|
||||
cfg,
|
||||
latent_image=None,
|
||||
start_step=None,
|
||||
last_step=None,
|
||||
force_full_denoise=False,
|
||||
denoise_mask=None,
|
||||
sigmas=None,
|
||||
callback=None,
|
||||
disable_pbar=False,
|
||||
seed=None,
|
||||
):
|
||||
if sigmas is None:
|
||||
sigmas = self.sigmas
|
||||
|
||||
from comfy.k_diffusion.sampling import BrownianTreeNoiseSampler
|
||||
|
||||
BrownianTreeNoiseSampler.set_global_sigmas(self.sigmas)
|
||||
|
||||
return super().sample(
|
||||
noise,
|
||||
positive,
|
||||
negative,
|
||||
cfg,
|
||||
latent_image,
|
||||
start_step,
|
||||
last_step,
|
||||
force_full_denoise,
|
||||
denoise_mask,
|
||||
sigmas,
|
||||
callback,
|
||||
disable_pbar,
|
||||
seed,
|
||||
)
|
||||
|
||||
hijack(mod, cls, HijackedKSampler)
|
||||
|
||||
|
||||
def hijack_browniannoisesampler(module, cls):
|
||||
mod = sys.modules[module]
|
||||
orig_sampler = getattr(mod, cls)
|
||||
if has_hijack(orig_sampler):
|
||||
return
|
||||
|
||||
class PCBrownianTreeNoiseSampler(orig_sampler):
|
||||
global_instance = None
|
||||
use_global_sigmas = False
|
||||
global_sigmas = None
|
||||
force_full_denoise = False
|
||||
|
||||
@classmethod
|
||||
def pc_reset(cls, use_global_sigmas=False, force_full_denoise=False):
|
||||
cls.global_instance = None
|
||||
cls.global_sigmas = None
|
||||
cls.use_global_sigmas = use_global_sigmas
|
||||
cls.force_full_denoise = force_full_denoise
|
||||
|
||||
@classmethod
|
||||
def set_global_sigmas(cls, sigmas):
|
||||
if cls.global_sigmas is None and cls.use_global_sigmas:
|
||||
cls.global_sigmas = (0 if cls.force_full_denoise else sigmas[sigmas > 0].min(), sigmas.max())
|
||||
log.info(
|
||||
"Initializing BrownianTreeNoiseSampler instance with global sigmas %s, %s",
|
||||
cls.global_sigmas,
|
||||
cls.force_full_denoise,
|
||||
)
|
||||
|
||||
def __init__(self, x, sigma_min, sigma_max, **kwargs):
|
||||
if self.global_sigmas is not None:
|
||||
sigma_min, sigma_max = self.global_sigmas
|
||||
if not self.global_instance:
|
||||
super().__init__(x, sigma_min, sigma_max, **kwargs)
|
||||
PCBrownianTreeNoiseSampler.global_instance = self
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
if self.global_instance and self != self.global_instance:
|
||||
return self.global_instance(*args, **kwargs)
|
||||
else:
|
||||
return super().__call__(*args, **kwargs)
|
||||
|
||||
hijack(mod, cls, PCBrownianTreeNoiseSampler)
|
||||
|
||||
|
||||
def do_hijack():
|
||||
hijack_browniannoisesampler("comfy.k_diffusion.sampling", "BrownianTreeNoiseSampler")
|
||||
hijack_sampler("comfy.sample", "sample", False)
|
||||
hijack_sampler("comfy.sample", "sample_custom", True)
|
||||
hijack_ksampler("comfy.samplers", "KSampler")
|
||||
@@ -1,48 +0,0 @@
|
||||
from .node_clip import control_to_clip_common
|
||||
from .node_lora import schedule_lora_common
|
||||
from ..parser import parse_prompt_schedules
|
||||
|
||||
|
||||
class PromptControlSimple:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"clip": ("CLIP",),
|
||||
"positive": ("STRING", {"multiline": True}),
|
||||
"negative": ("STRING", {"multiline": True}),
|
||||
},
|
||||
"optional": {
|
||||
"tags": ("STRING", {"default": ""}),
|
||||
"start": ("FLOAT", {"min": 0.0, "max": 1.0, "step": 0.1, "default": 0.0}),
|
||||
"end": ("FLOAT", {"min": 0.0, "max": 1.0, "step": 0.1, "default": 1.0}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CONDITIONING", "CONDITIONING", "MODEL", "CONDITIONING", "CONDITIONING")
|
||||
RETURN_NAMES = ("model", "positive", "negative", "model_filtered", "pos_filtered", "neg_filtered")
|
||||
CATEGORY = "promptcontrol/_legacy"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, clip, positive, negative, tags="", start=0.0, end=1.0):
|
||||
lora_cache = {}
|
||||
cond_cache = {}
|
||||
pos_sched = parse_prompt_schedules(positive)
|
||||
pos_cond = pos_filtered = control_to_clip_common(clip, pos_sched, lora_cache, cond_cache)
|
||||
|
||||
neg_sched = parse_prompt_schedules(negative)
|
||||
neg_cond = neg_filtered = control_to_clip_common(clip, neg_sched, lora_cache, cond_cache)
|
||||
|
||||
new_model = model_filtered = schedule_lora_common(model, pos_sched, lora_cache)
|
||||
|
||||
if [tags.strip(), start, end] != ["", 0.0, 1.0]:
|
||||
pos_filtered = control_to_clip_common(
|
||||
clip, pos_sched.with_filters(tags, start, end), lora_cache, cond_cache
|
||||
)
|
||||
neg_filtered = control_to_clip_common(
|
||||
clip, neg_sched.with_filters(tags, start, end), lora_cache, cond_cache
|
||||
)
|
||||
model_filtered = schedule_lora_common(model, pos_sched.with_filters(tags, start, end), lora_cache)
|
||||
|
||||
return (new_model, pos_cond, neg_cond, model_filtered, pos_filtered, neg_filtered)
|
||||
@@ -1,701 +0,0 @@
|
||||
import logging
|
||||
import re
|
||||
import torch
|
||||
from ..parser import parse_prompt_schedules, parse_cuts
|
||||
from .utils import Timer, equalize, apply_loras_from_spec
|
||||
from ..utils import safe_float, get_function, parse_floats # non-legacy
|
||||
from .perp_weight import perp_encode
|
||||
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
|
||||
from node_helpers import conditioning_set_values
|
||||
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
try:
|
||||
from custom_nodes.ComfyUI_ADV_CLIP_emb.adv_encode import (
|
||||
advanced_encode_from_tokens,
|
||||
encode_token_weights_l,
|
||||
encode_token_weights_g,
|
||||
prepareXL,
|
||||
encode_token_weights,
|
||||
)
|
||||
|
||||
have_advanced_encode = True
|
||||
AVAILABLE_STYLES = ["comfy", "A1111", "compel", "comfy++", "down_weight"]
|
||||
AVAILABLE_NORMALIZATIONS = ["none", "mean", "length", "length+mean"]
|
||||
except ImportError:
|
||||
have_advanced_encode = False
|
||||
AVAILABLE_STYLES = ["comfy"]
|
||||
AVAILABLE_NORMALIZATIONS = ["none"]
|
||||
|
||||
try:
|
||||
from custom_nodes.Vector_Sculptor_ComfyUI.nodes import vector_sculptor_tokens
|
||||
|
||||
can_sculpt = True
|
||||
log.info("Vector sculptor extension detected, can use SCULPT()")
|
||||
except ImportError:
|
||||
can_sculpt = False
|
||||
|
||||
|
||||
AVAILABLE_STYLES.append("perp")
|
||||
log.info("Use STYLE(weight_interpretation, normalization) at the start of a prompt to use advanced encodings")
|
||||
log.info("Weight interpretations available: %s", ",".join(AVAILABLE_STYLES))
|
||||
log.info("Normalization types available: %s", ",".join(AVAILABLE_NORMALIZATIONS))
|
||||
|
||||
|
||||
def linear_interpolate_cond(
|
||||
start, end, from_step=0.0, to_step=1.0, step=0.1, start_at=None, end_at=None, prompt_start="N/A", prompt_end="N/A"
|
||||
):
|
||||
count = min(len(start), len(end))
|
||||
if len(start) != len(end):
|
||||
log.info(
|
||||
"Length of conds to interpolate does not match (start=%s != end=%s), interpolating up to %s.",
|
||||
len(start),
|
||||
len(end),
|
||||
count,
|
||||
)
|
||||
|
||||
all_res = []
|
||||
for idx in range(count):
|
||||
res = []
|
||||
from_cond, to_cond = equalize(start[idx][0], end[idx][0])
|
||||
from_pooled = start[idx][1].get("pooled_output")
|
||||
to_pooled = end[idx][1].get("pooled_output")
|
||||
start_at = start_at if start_at is not None else from_step
|
||||
end_at = end_at if end_at is not None else to_step
|
||||
total_steps = int(round((to_step - from_step) / step, 0))
|
||||
num_steps = int(round((end_at - from_step) / step, 0))
|
||||
start_on = int(round((start_at - from_step) / step, 0))
|
||||
start_pct = start_at
|
||||
log.debug(
|
||||
f"interpolate_cond {idx=} {from_step=} {to_step=} {start_at=} {end_at=} {total_steps=} {num_steps=} {start_on=} {step=}"
|
||||
)
|
||||
x = 1 / (total_steps + 1)
|
||||
for s in range(start_on, num_steps):
|
||||
factor = round((s + 1) * x, 2)
|
||||
new_cond = from_cond + (to_cond - from_cond) * factor
|
||||
if from_pooled is not None and to_pooled is not None:
|
||||
from_pooled, to_pooled = equalize(from_pooled, to_pooled)
|
||||
new_pooled = from_pooled + (to_pooled - from_pooled) * factor
|
||||
elif from_pooled is not None:
|
||||
new_pooled = from_pooled
|
||||
|
||||
n = [new_cond, start[idx][1].copy()]
|
||||
if new_pooled is not None:
|
||||
n[1]["pooled_output"] = new_pooled
|
||||
n[1]["start_percent"] = round(start_pct, 2)
|
||||
n[1]["end_percent"] = min(round((start_pct + step), 2), 1.0)
|
||||
start_pct += step
|
||||
start_pct = round(start_pct, 2)
|
||||
if prompt_start:
|
||||
n[1]["prompt"] = f"linear:{round(1.0 - factor, 2)} / {factor}"
|
||||
log.debug(
|
||||
"Interpolating at step %s with factor %s (%s, %s)...",
|
||||
s,
|
||||
factor,
|
||||
n[1]["start_percent"],
|
||||
n[1]["end_percent"],
|
||||
)
|
||||
res.append(n)
|
||||
if res:
|
||||
res[-1][1]["end_percent"] = round(end_at, 2)
|
||||
all_res.extend(res)
|
||||
return all_res
|
||||
|
||||
|
||||
def get_control_points(schedule, steps, encoder):
|
||||
assert len(steps) > 1
|
||||
new_steps = set(steps)
|
||||
|
||||
for step in (s[0] for s in schedule if s[0] >= steps[0] and s[0] <= steps[-1]):
|
||||
new_steps.add(step)
|
||||
control_points = [(s, encoder(schedule.at_step(s)[1])) for s in new_steps]
|
||||
log.debug("Actual control points for interpolation: %s (from %s)", new_steps, steps)
|
||||
return sorted(control_points, key=lambda x: x[0])
|
||||
|
||||
|
||||
def linear_interpolator(control_points, step, start_pct, end_pct):
|
||||
o_start, start = control_points[0]
|
||||
o_end, _ = control_points[-1]
|
||||
t_start = o_start
|
||||
conds = []
|
||||
for t_end, end in control_points[1:]:
|
||||
if t_start < start_pct:
|
||||
t_start, start = t_end, end
|
||||
continue
|
||||
if t_start >= end_pct:
|
||||
break
|
||||
cs = linear_interpolate_cond(start, end, o_start, o_end, step, start_at=t_start, end_at=end_pct)
|
||||
if cs:
|
||||
conds.extend(cs)
|
||||
else:
|
||||
break
|
||||
t_start = t_end
|
||||
start = end
|
||||
return conds
|
||||
|
||||
|
||||
class ScheduleToCond:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {"clip": ("CLIP",), "prompt_schedule": ("PROMPT_SCHEDULE",)},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "promptcontrol/_legacy"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, clip, prompt_schedule):
|
||||
with Timer("ScheduleToCond"):
|
||||
r = (control_to_clip_common(clip, prompt_schedule),)
|
||||
return r
|
||||
|
||||
|
||||
class EditableCLIPEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"clip": ("CLIP",),
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
},
|
||||
"optional": {"filter_tags": ("STRING", {"default": ""})},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
CATEGORY = "promptcontrol/_donotuse"
|
||||
FUNCTION = "parse"
|
||||
|
||||
def parse(self, clip, text, filter_tags=""):
|
||||
parsed = parse_prompt_schedules(text).with_filters(filter_tags)
|
||||
return (control_to_clip_common(clip, parsed),)
|
||||
|
||||
|
||||
def get_sdxl(text, defaults):
|
||||
# Defaults fail to parse and get looked up from the defaults dict
|
||||
text, sdxl = get_function(text, "SDXL", ["none", "none", "none"])
|
||||
if not sdxl:
|
||||
return text, {}
|
||||
args = sdxl[0]
|
||||
d = defaults
|
||||
w, h = parse_floats(args[0], [d.get("sdxl_width", 1024), d.get("sdxl_height", 1024)], split_re="\\s+")
|
||||
tw, th = parse_floats(args[1], [d.get("sdxl_twidth", 1024), d.get("sdxl_theight", 1024)], split_re="\\s+")
|
||||
cropw, croph = parse_floats(args[2], [d.get("sdxl_cwidth", 0), d.get("sdxl_cheight", 0)], split_re="\\s+")
|
||||
|
||||
opts = {
|
||||
"width": int(w),
|
||||
"height": int(h),
|
||||
"target_width": int(tw),
|
||||
"target_height": int(th),
|
||||
"crop_w": int(cropw),
|
||||
"crop_h": int(croph),
|
||||
}
|
||||
return text, opts
|
||||
|
||||
|
||||
def get_style(text, default_style="comfy", default_normalization="none"):
|
||||
text, styles = get_function(text, "STYLE", [default_style, default_normalization])
|
||||
if not styles:
|
||||
return default_style, default_normalization, text
|
||||
style, normalization = styles[0]
|
||||
style = style.strip()
|
||||
normalization = normalization.strip()
|
||||
if style not in AVAILABLE_STYLES:
|
||||
log.warning("Unrecognized prompt style: %s. Using %s", style, default_style)
|
||||
style = default_style
|
||||
|
||||
if normalization not in AVAILABLE_NORMALIZATIONS:
|
||||
log.warning("Unrecognized prompt normalization: %s. Using %s", normalization, default_normalization)
|
||||
normalization = default_normalization
|
||||
|
||||
return style, normalization, text
|
||||
|
||||
|
||||
def encode_regions(clip, tokens, regions, weight_interpretation="comfy", token_normalization="none"):
|
||||
from custom_nodes.ComfyUI_Cutoff.cutoff import CLIPSetRegion, finalize_clip_regions
|
||||
|
||||
clip_regions = {
|
||||
"clip": clip,
|
||||
"base_tokens": tokens,
|
||||
"regions": [],
|
||||
"targets": [],
|
||||
"weights": [],
|
||||
}
|
||||
|
||||
strict_mask = 1.0
|
||||
start_from_masked = 1.0
|
||||
mask_token = ""
|
||||
|
||||
for region in regions:
|
||||
region_text, target_text, w, sm, sfm, mt = region
|
||||
if w is not None:
|
||||
w = safe_float(w, 0)
|
||||
else:
|
||||
w = 1.0
|
||||
if sm is not None:
|
||||
strict_mask = safe_float(sm, 1.0)
|
||||
if sfm is not None:
|
||||
start_from_masked = safe_float(sfm, 1.0)
|
||||
if mt is not None:
|
||||
mask_token = mt
|
||||
log.info("Region: text %s, target %s, weight %s", region_text.strip(), target_text.strip(), w)
|
||||
(clip_regions,) = CLIPSetRegion.add_clip_region(None, clip_regions, region_text, target_text, w)
|
||||
log.info("Regions: mask_token=%s strict_mask=%s start_from_masked=%s", mask_token, strict_mask, start_from_masked)
|
||||
|
||||
(r,) = finalize_clip_regions(
|
||||
clip_regions, mask_token, strict_mask, start_from_masked, token_normalization, weight_interpretation
|
||||
)
|
||||
cond, pooled = r[0][0], r[0][1].get("pooled_output")
|
||||
return cond, pooled
|
||||
|
||||
|
||||
SHUFFLE_GEN = torch.Generator(device="cpu")
|
||||
|
||||
|
||||
def shuffle_chunk(shuffle, c):
|
||||
func, shuffle = shuffle
|
||||
shuffle_count = int(safe_float(shuffle[0], 0))
|
||||
_, separator, joiner = shuffle
|
||||
if separator == "default":
|
||||
separator = ","
|
||||
|
||||
if not separator:
|
||||
separator = ","
|
||||
|
||||
joiner = {
|
||||
"default": ",",
|
||||
"separator": separator,
|
||||
}.get(joiner, joiner)
|
||||
|
||||
log.info("%s arg=%s sep=%s join=%s", func, shuffle_count, separator, joiner)
|
||||
separated = c.split(separator)
|
||||
if func == "SHIFT":
|
||||
shuffle_count = shuffle_count % len(separated)
|
||||
permutation = separated[shuffle_count:] + separated[:shuffle_count]
|
||||
elif func == "SHUFFLE":
|
||||
SHUFFLE_GEN.manual_seed(shuffle_count)
|
||||
permutation = [separated[i] for i in torch.randperm(len(separated), generator=SHUFFLE_GEN)]
|
||||
else:
|
||||
# ??? should never get here
|
||||
permutation = separated
|
||||
|
||||
permutation = [p for p in permutation if p.strip()]
|
||||
if permutation != separated:
|
||||
c = joiner.join(permutation)
|
||||
return c
|
||||
|
||||
|
||||
def fix_word_ids(tokens):
|
||||
"""Fix word indexes. Tokenizing separately (when BREAKs exist) causes the indexes to restart which causes problems with some weighting algorithms that rely on them"""
|
||||
for key in tokens:
|
||||
max_idx = 0
|
||||
for group in range(len(tokens[key])):
|
||||
for i, token in enumerate(tokens[key][group]):
|
||||
if len(token) < 3:
|
||||
# No need to fix ids when they don't exist
|
||||
return tokens
|
||||
# Ignore zeros, they represent the padding token
|
||||
if token[2] != 0 and token[2] < max_idx:
|
||||
tokens[key][group][i] = (token[0], token[1], token[2] + max_idx)
|
||||
max_idx = max(max_idx, max(x for _, _, x in tokens[key][group]))
|
||||
return tokens
|
||||
|
||||
|
||||
def encode_prompt(clip, text, default_style="comfy", default_normalization="none"):
|
||||
style, normalization, text = get_style(text, default_style, default_normalization)
|
||||
sculpts = []
|
||||
if can_sculpt:
|
||||
text, sculpts = get_function(text, "SCULPT", ["1.0", "forward", "none"])
|
||||
text, regions = parse_cuts(text)
|
||||
# defaults=None means there is no argument parsing at all
|
||||
text, l_prompts = get_function(text, "CLIP_L", defaults=None)
|
||||
chunks = re.split(r"\bBREAK\b", text)
|
||||
token_chunks = []
|
||||
need_word_ids = len(regions) > 0 or (have_advanced_encode and style != "perp")
|
||||
for c in chunks:
|
||||
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"], return_func_name=True)
|
||||
r = c
|
||||
for s in shuffles:
|
||||
r = shuffle_chunk(s, r)
|
||||
if r != c:
|
||||
log.info("Shuffled prompt chunk to %s", r)
|
||||
c = r
|
||||
if sculpts:
|
||||
w, method, norm = sculpts[0]
|
||||
log.info("Using vector sculptor with method=%s norm=%s w=%s", method, norm, w)
|
||||
w = safe_float(w, 1.0)
|
||||
t = vector_sculptor_tokens(clip, c, method, norm, w)
|
||||
else:
|
||||
# Tokenizer returns padded results
|
||||
t = clip.tokenize(c, return_word_ids=need_word_ids)
|
||||
token_chunks.append(t)
|
||||
tokens = token_chunks[0]
|
||||
|
||||
for key in tokens:
|
||||
for c in token_chunks[1:]:
|
||||
tokens[key].extend(c[key])
|
||||
|
||||
# Non-SDXL has only "l"
|
||||
if "g" in tokens and l_prompts:
|
||||
text_l = " ".join(l_prompts)
|
||||
log.info("Encoded SDXL CLIP_L prompt: %s", text_l)
|
||||
tokens["l"] = clip.tokenize(text_l, return_word_ids=need_word_ids)["l"]
|
||||
|
||||
if "g" in tokens and "l" in tokens and len(tokens["l"]) != len(tokens["g"]):
|
||||
empty = clip.tokenize("", return_word_ids=need_word_ids)
|
||||
while len(tokens["l"]) < len(tokens["g"]):
|
||||
tokens["l"] += empty["l"]
|
||||
while len(tokens["l"]) > len(tokens["g"]):
|
||||
tokens["g"] += empty["g"]
|
||||
|
||||
tokens = fix_word_ids(tokens)
|
||||
|
||||
if len(regions) > 0:
|
||||
return encode_regions(clip, tokens, regions, style, normalization)
|
||||
|
||||
if style == "perp":
|
||||
if normalization != "none":
|
||||
log.warning("Normalization is not supported with perp style weighting. Ignored '%s'", normalization)
|
||||
return perp_encode(clip, tokens)
|
||||
|
||||
if "t5xxl" not in tokens and have_advanced_encode and not sculpts:
|
||||
if "g" in tokens:
|
||||
embs_l = None
|
||||
embs_g = None
|
||||
pooled = None
|
||||
if "l" in tokens:
|
||||
embs_l, _ = advanced_encode_from_tokens(
|
||||
tokens["l"],
|
||||
normalization,
|
||||
style,
|
||||
lambda x: encode_token_weights(clip, x, encode_token_weights_l),
|
||||
return_pooled=False,
|
||||
)
|
||||
if "g" in tokens:
|
||||
embs_g, pooled = advanced_encode_from_tokens(
|
||||
tokens["g"],
|
||||
normalization,
|
||||
style,
|
||||
lambda x: encode_token_weights(clip, x, encode_token_weights_g),
|
||||
return_pooled=True,
|
||||
apply_to_pooled=False,
|
||||
)
|
||||
# Hardcoded clip_balance
|
||||
return prepareXL(embs_l, embs_g, pooled, 0.5)
|
||||
return advanced_encode_from_tokens(
|
||||
tokens["l"],
|
||||
normalization,
|
||||
style,
|
||||
lambda x: clip.encode_from_tokens({"l": x}, return_pooled=True),
|
||||
return_pooled=True,
|
||||
apply_to_pooled=True,
|
||||
)
|
||||
else:
|
||||
return clip.encode_from_tokens(tokens, return_pooled=True)
|
||||
|
||||
|
||||
def get_area(text):
|
||||
text, areas = get_function(text, "AREA", ["0 1", "0 1", "1"])
|
||||
if not areas:
|
||||
return text, None
|
||||
|
||||
args = areas[0]
|
||||
x, w = parse_floats(args[0], [0.0, 1.0], split_re="\\s+")
|
||||
y, h = parse_floats(args[1], [0.0, 1.0], split_re="\\s+")
|
||||
weight = safe_float(args[2], 1.0)
|
||||
|
||||
def is_pct(f):
|
||||
return f >= 0.0 and f <= 1.0
|
||||
|
||||
def is_pixel(f):
|
||||
return f == 0 or f > 1
|
||||
|
||||
if all(is_pct(v) for v in [h, w, y, x]):
|
||||
area = ("percentage", h, w, y, x)
|
||||
elif all(is_pixel(v) for v in [h, w, y, x]):
|
||||
area = (int(h) // 8, int(w) // 8, int(y) // 8, int(x) // 8)
|
||||
else:
|
||||
raise Exception(
|
||||
f"AREA specified with invalid size {x} {w}, {h} {y}. They must either all be percentages between 0 and 1 or positive integer pixel values excluding 1"
|
||||
)
|
||||
|
||||
return text, (area, weight)
|
||||
|
||||
|
||||
def get_mask_size(text, defaults):
|
||||
text, sizes = get_function(text, "MASK_SIZE", ["512", "512"])
|
||||
if not sizes:
|
||||
return text, (defaults.get("mask_width", 512), defaults.get("mask_height", 512))
|
||||
w, h = sizes[0]
|
||||
return text, (int(w), int(h))
|
||||
|
||||
|
||||
def make_mask(args, size, weight):
|
||||
x1, x2 = parse_floats(args[0], [0.0, 1.0], split_re="\\s+")
|
||||
y1, y2 = parse_floats(args[1], [0.0, 1.0], split_re="\\s+")
|
||||
|
||||
def is_pct(f):
|
||||
return f >= 0.0 and f <= 1.0
|
||||
|
||||
def is_pixel(f):
|
||||
return f == 0 or f > 1
|
||||
|
||||
if all(is_pct(v) for v in [x1, x2, y1, y2]):
|
||||
w, h = size
|
||||
xs = int(w * x1), int(w * x2)
|
||||
ys = int(h * y1), int(h * y2)
|
||||
elif all(is_pixel(v) for v in [x1, x2, y1, y2]):
|
||||
w, h = size
|
||||
xs = int(x1), int(x2)
|
||||
ys = int(y1), int(y2)
|
||||
else:
|
||||
raise Exception(
|
||||
f"MASK specified with invalid size {x1} {x2}, {y1} {y2}. They must either all be percentages between 0 and 1 or positive integer pixel values excluding 1"
|
||||
)
|
||||
|
||||
mask = torch.full((h, w), 0, dtype=torch.float32, device="cpu")
|
||||
mask[ys[0] : ys[1], xs[0] : xs[1]] = weight
|
||||
mask = mask.unsqueeze(0)
|
||||
log.info("Mask xs=%s, ys=%s, shape=%s, weight=%s", xs, ys, mask.shape, weight)
|
||||
return mask
|
||||
|
||||
|
||||
def get_mask(text, size, input_masks):
|
||||
"""Parse MASK(x1 x2, y1 y2, weight), IMASK(i, weight) and FEATHER(left top right bottom)"""
|
||||
# TODO: combine multiple masks
|
||||
text, masks = get_function(text, "MASK", ["0 1", "0 1", "1", "multiply"])
|
||||
text, imasks = get_function(text, "IMASK", ["0", "1", "multiply"])
|
||||
text, feathers = get_function(text, "FEATHER", ["0 0 0 0"])
|
||||
text, maskw = get_function(text, "MASKW", ["1.0"])
|
||||
if not masks and not imasks:
|
||||
return text, None, None
|
||||
|
||||
def feather(f, mask):
|
||||
l, t, r, b, *_ = [int(x) for x in parse_floats(f[0], [0, 0, 0, 0], split_re="\\s+")]
|
||||
mask = FeatherMask().feather(mask, l, t, r, b)[0]
|
||||
log.info("FeatherMask l=%s, t=%s, r=%s, b=%s", l, t, r, b)
|
||||
return mask
|
||||
|
||||
mask = None
|
||||
totalweight = 1.0
|
||||
if maskw:
|
||||
totalweight = safe_float(maskw[0][0], 1.0)
|
||||
i = 0
|
||||
for m in masks:
|
||||
weight = safe_float(m[2], 1.0)
|
||||
op = m[3]
|
||||
nextmask = make_mask(m, size, weight)
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
log.info("MaskComposite op=%s", op)
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
for idx, w, op in imasks:
|
||||
idx = int(safe_float(idx, 0.0))
|
||||
w = safe_float(w, 1.0)
|
||||
if len(input_masks) < idx + 1:
|
||||
log.warn("IMASK index %s not found, ignoring...", idx)
|
||||
continue
|
||||
nextmask = input_masks[idx] * w
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
mask = MaskComposite().combine(mask, nextmask, 0, 0, op)[0]
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
# apply leftover FEATHER() specs to the whole
|
||||
for f in feathers[i:]:
|
||||
mask = feather(f, mask)
|
||||
|
||||
return text, mask, totalweight
|
||||
|
||||
|
||||
def get_noise(text):
|
||||
text, noises = get_function(
|
||||
text,
|
||||
"NOISE",
|
||||
["0.0", "none"],
|
||||
)
|
||||
if not noises:
|
||||
return text, None, None
|
||||
w = 0
|
||||
# Only take seed from first noise spec, for simplicity
|
||||
seed = safe_float(noises[0][1], "none")
|
||||
if seed == "none":
|
||||
gen = None
|
||||
else:
|
||||
gen = torch.Generator()
|
||||
gen.manual_seed(int(seed))
|
||||
for n in noises:
|
||||
w += safe_float(n[0], 0.0)
|
||||
return text, max(min(w, 1.0), 0.0), gen
|
||||
|
||||
|
||||
def apply_noise(cond, weight, gen):
|
||||
if cond is None or not weight:
|
||||
return cond
|
||||
|
||||
n = torch.randn(cond.size(), generator=gen).to(cond)
|
||||
|
||||
return cond * (1 - weight) + n * weight
|
||||
|
||||
|
||||
def do_encode(clip, text, defaults, masks):
|
||||
# First style modifier applies to ANDed prompts too unless overridden
|
||||
style, normalization, text = get_style(text)
|
||||
text, mask_size = get_mask_size(text, defaults)
|
||||
|
||||
# Don't sum ANDs if this is in prompt
|
||||
alt_method = "COMFYAND()" in text
|
||||
text = text.replace("COMFYAND()", "")
|
||||
|
||||
prompts = [p.strip() for p in re.split(r"\bAND\b", text)]
|
||||
|
||||
p, sdxl_opts = get_sdxl(prompts[0], defaults)
|
||||
prompts[0] = p
|
||||
|
||||
def weight(t):
|
||||
opts = {}
|
||||
m = re.search(r":(-?\d\.?\d*)(![A-Za-z]+)?$", t)
|
||||
if not m:
|
||||
return (1.0, opts, t)
|
||||
w = float(m[1])
|
||||
tag = m[2]
|
||||
t = t[: m.span()[0]]
|
||||
if tag == "!noscale":
|
||||
opts["scale"] = 1
|
||||
|
||||
return w, opts, t
|
||||
|
||||
conds = []
|
||||
res = []
|
||||
scale = sum(abs(weight(p)[0]) for p in prompts if not ("AREA(" in p or "MASK(" in p))
|
||||
for prompt in prompts:
|
||||
prompt, mask, mask_weight = get_mask(prompt, mask_size, masks)
|
||||
w, opts, prompt = weight(prompt)
|
||||
text, noise_w, generator = get_noise(text)
|
||||
if not w:
|
||||
continue
|
||||
prompt, area = get_area(prompt)
|
||||
prompt, local_sdxl_opts = get_sdxl(prompt, defaults)
|
||||
cond, pooled = encode_prompt(clip, prompt, style, normalization)
|
||||
cond = apply_noise(cond, noise_w, generator)
|
||||
pooled = apply_noise(pooled, noise_w, generator)
|
||||
|
||||
settings = {"prompt": prompt}
|
||||
if alt_method:
|
||||
settings["strength"] = w
|
||||
settings.update(sdxl_opts)
|
||||
settings.update(local_sdxl_opts)
|
||||
if area:
|
||||
settings["area"] = area[0]
|
||||
settings["strength"] = area[1]
|
||||
settings["set_area_to_bounds"] = False
|
||||
if mask is not None:
|
||||
settings["mask"] = mask
|
||||
settings["mask_strength"] = mask_weight
|
||||
|
||||
if mask is not None or area or alt_method or local_sdxl_opts:
|
||||
if pooled is not None:
|
||||
settings["pooled_output"] = pooled
|
||||
conds.append([cond, settings])
|
||||
else:
|
||||
s = opts.get("scale", scale)
|
||||
res.append((cond, pooled, w / s))
|
||||
|
||||
sumconds = [r[0] * r[2] for r in res]
|
||||
pooleds = [r[1] for r in res if r[1] is not None]
|
||||
|
||||
if len(res) > 0:
|
||||
opts = sdxl_opts
|
||||
if pooleds:
|
||||
opts["pooled_output"] = sum(equalize(*pooleds))
|
||||
sumcond = sum(equalize(*sumconds))
|
||||
conds.append([sumcond, opts])
|
||||
return conds
|
||||
|
||||
|
||||
def debug_conds(conds):
|
||||
r = []
|
||||
for i, c in enumerate(conds):
|
||||
x = c[1].copy()
|
||||
if "pooled_output" in x:
|
||||
del x["pooled_output"]
|
||||
r.append((i, x))
|
||||
return r
|
||||
|
||||
|
||||
def control_to_clip_common(clip, schedules, lora_cache=None, cond_cache=None):
|
||||
orig_clip = clip.clone()
|
||||
current_loras = {}
|
||||
if lora_cache is None:
|
||||
lora_cache = {}
|
||||
start_pct = 0.0
|
||||
conds = []
|
||||
cond_cache = cond_cache if cond_cache is not None else {}
|
||||
|
||||
def c_str(c):
|
||||
r = [c["prompt"]]
|
||||
loras = c["loras"]
|
||||
for k in sorted(loras.keys()):
|
||||
r.append(k)
|
||||
r.append(loras[k]["weight_clip"])
|
||||
for lbw, val in loras[k].get("lbw", {}).items():
|
||||
r.append(lbw)
|
||||
r.append(val)
|
||||
return "".join(str(i) for i in r)
|
||||
|
||||
def encode(c):
|
||||
nonlocal clip
|
||||
nonlocal current_loras
|
||||
prompt = c["prompt"]
|
||||
loras = c["loras"]
|
||||
cachekey = c_str(c)
|
||||
cond = cond_cache.get(cachekey)
|
||||
if cond is None:
|
||||
if loras != current_loras:
|
||||
_, clip = apply_loras_from_spec(loras, clip=orig_clip, cache=lora_cache, applied_loras=current_loras)
|
||||
current_loras = loras
|
||||
cond_cache[cachekey] = do_encode(clip, prompt, schedules.defaults, schedules.masks)
|
||||
return cond_cache[cachekey]
|
||||
|
||||
for end_pct, c in schedules:
|
||||
interpolations = [
|
||||
i
|
||||
for i in schedules.interpolations
|
||||
if (start_pct >= i[0][0] and start_pct < i[0][-1]) or (end_pct > i[0][0] and start_pct < i[0][-1])
|
||||
]
|
||||
new_start_pct = start_pct
|
||||
if interpolations:
|
||||
min_step = min(i[1] for i in interpolations)
|
||||
for i in interpolations:
|
||||
control_points, _ = i
|
||||
interpolation_end_pct = min(control_points[-1], end_pct)
|
||||
interpolation_start_pct = max(control_points[0], start_pct)
|
||||
|
||||
control_points = get_control_points(schedules, control_points, encode)
|
||||
cs = linear_interpolator(control_points, min_step, interpolation_start_pct, interpolation_end_pct)
|
||||
conds.extend(cs)
|
||||
new_start_pct = max(new_start_pct, interpolation_end_pct)
|
||||
start_pct = new_start_pct
|
||||
|
||||
if start_pct < end_pct:
|
||||
cond = encode(c)
|
||||
# Node functions return lists of cond
|
||||
cond = conditioning_set_values(
|
||||
cond, {"start_percent": round(start_pct, 2), "end_percent": round(end_pct, 2), "prompt": c["prompt"]}
|
||||
)
|
||||
conds.extend(cond)
|
||||
|
||||
start_pct = end_pct
|
||||
log.debug("Conds at the end: %s", debug_conds(conds))
|
||||
|
||||
log.debug("Final cond info: %s", debug_conds(conds))
|
||||
return conds
|
||||
@@ -1,247 +0,0 @@
|
||||
import logging
|
||||
import torch
|
||||
|
||||
from .utils import unpatch_model, clone_model, set_callback, apply_loras_from_spec
|
||||
from ..parser import parse_prompt_schedules
|
||||
from .hijack import do_hijack
|
||||
from comfy.samplers import CFGGuider
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def apply_lora_for_step(schedules, step, total_steps, state, original_model, lora_cache, patch=True):
|
||||
# zero-indexed steps, 0 = first step, but schedules are 1-indexed
|
||||
sched = schedules.at_step(step + 1, total_steps)
|
||||
lora_spec = sched[1]["loras"]
|
||||
|
||||
if state["applied_loras"] != lora_spec:
|
||||
log.debug("At step %s, applying lora_spec %s", step, lora_spec)
|
||||
m, _ = apply_loras_from_spec(
|
||||
lora_spec,
|
||||
model=state["model"],
|
||||
orig_model=original_model,
|
||||
cache=lora_cache,
|
||||
patch=patch,
|
||||
applied_loras=state["applied_loras"],
|
||||
)
|
||||
state["model"] = m
|
||||
state["applied_loras"] = lora_spec
|
||||
|
||||
|
||||
def schedule_lora_common(model, schedules, lora_cache=None):
|
||||
do_hijack()
|
||||
orig_model = clone_model(model)
|
||||
orig_model.model_options["pc_schedules"] = schedules
|
||||
|
||||
if lora_cache is None:
|
||||
lora_cache = {}
|
||||
|
||||
def sampler_cb(orig_sampler, is_custom, *args, **kwargs):
|
||||
split_sampling = args[0].model_options.get("pc_split_sampling")
|
||||
state = {}
|
||||
if is_custom:
|
||||
steps = len(args[4])
|
||||
log.info(
|
||||
"SamplerCustom detected, number of steps not available. LoRA schedules will be calculated based on the number of sigmas (%s)",
|
||||
steps,
|
||||
)
|
||||
else:
|
||||
log.debug("Normal sampler detected, using steps from parameter")
|
||||
steps = args[2]
|
||||
start_step = kwargs.get("start_step") or 0
|
||||
# The model patcher may change if LoRAs are applied
|
||||
state["model"] = args[0]
|
||||
state["applied_loras"] = {}
|
||||
|
||||
orig_cb = kwargs["callback"]
|
||||
|
||||
def step_callback(*args, **kwargs):
|
||||
current_step = args[0] + start_step
|
||||
apply_lora_for_step(schedules, current_step, steps, state, orig_model, lora_cache, patch=True)
|
||||
if orig_cb:
|
||||
return orig_cb(*args, **kwargs)
|
||||
|
||||
kwargs["callback"] = step_callback
|
||||
|
||||
apply_lora_for_step(schedules, start_step, steps, state, orig_model, lora_cache, patch=True)
|
||||
|
||||
def filter_conds(conds, t, start_t, end_t):
|
||||
r = []
|
||||
for c in conds:
|
||||
x = c[1].copy()
|
||||
start_at = round(x["start_percent"], 2)
|
||||
end_at = round(x["end_percent"], 2)
|
||||
# Take any cond that has any effect before end_t, since the percentages may not perfectly match
|
||||
if end_t > start_at and end_t <= end_at:
|
||||
del x["start_percent"]
|
||||
del x["end_percent"]
|
||||
r.append([c[0].clone(), x])
|
||||
else:
|
||||
log.debug("Rejecting cond (%s, %s) between (%s, %s)", start_at, end_at, start_t, end_t)
|
||||
if len(r) == 0:
|
||||
log.error("No %s conds between (%s, %s); Try adjusting your steps", t, start_t, end_t)
|
||||
return r
|
||||
|
||||
def get_steps(conds):
|
||||
for c in conds:
|
||||
yield round(c[1].get("end_percent", 0), 2)
|
||||
|
||||
if split_sampling:
|
||||
actual_end_step = kwargs["last_step"] or steps
|
||||
first_step = True
|
||||
s = args[8]
|
||||
all_steps = sorted(set(int(steps * i) for i in [1.0] + list(get_steps(args[6])) + list(get_steps(args[7]))))
|
||||
for end_step in all_steps:
|
||||
if end_step <= start_step:
|
||||
continue
|
||||
start_t = round(start_step / steps, 2)
|
||||
end_t = round(end_step / steps, 2)
|
||||
new_kwargs = kwargs.copy()
|
||||
new_args = list(args)
|
||||
new_args[0] = state["model"]
|
||||
new_args[6] = filter_conds(new_args[6], "positive", start_t, end_t)
|
||||
new_args[7] = filter_conds(new_args[7], "negative", start_t, end_t)
|
||||
new_args[8] = s
|
||||
log.info("Sampling from %s to %s (total: %s)", start_step, end_step, actual_end_step)
|
||||
new_kwargs["start_step"] = start_step
|
||||
new_kwargs["last_step"] = end_step
|
||||
if end_step >= min(steps, actual_end_step):
|
||||
new_kwargs["force_full_denoise"] = kwargs["force_full_denoise"]
|
||||
else:
|
||||
new_kwargs["force_full_denoise"] = False
|
||||
|
||||
if not first_step:
|
||||
# disable_noise apparently does nothing currently, we need to override noise in args
|
||||
new_kwargs["disable_noise"] = True
|
||||
new_args[1] = torch.zeros_like(s)
|
||||
|
||||
s = orig_sampler(*new_args, **new_kwargs)
|
||||
start_step = end_step
|
||||
first_step = False
|
||||
else:
|
||||
args = list(args)
|
||||
args[0] = state["model"]
|
||||
s = orig_sampler(*args, **kwargs)
|
||||
|
||||
unpatch_model(state["model"])
|
||||
|
||||
return s
|
||||
|
||||
set_callback(orig_model, sampler_cb)
|
||||
|
||||
return orig_model
|
||||
|
||||
|
||||
class PCWrapGuider:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"guider": ("GUIDER",),
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "promptcontrol/_legacy"
|
||||
FUNCTION = "apply"
|
||||
RETURN_TYPES = ("GUIDER",)
|
||||
|
||||
def apply(self, guider):
|
||||
return (PCGuider(guider),)
|
||||
|
||||
|
||||
class PCGuider(CFGGuider):
|
||||
def __init__(self, original_guider):
|
||||
if "pc_schedules" not in original_guider.model_patcher.model_options:
|
||||
raise ValueError(
|
||||
"The guider passed to PCWrapGuider must contain a Model that has schedules applied. Use ScheduleToModel"
|
||||
)
|
||||
self.schedules = original_guider.model_patcher.model_options["pc_schedules"]
|
||||
self.guider = original_guider
|
||||
self.lora_cache = {}
|
||||
# sets self.model_patcher
|
||||
super().__init__(original_guider.model_patcher)
|
||||
|
||||
def sample(self, *args, **kwargs):
|
||||
orig_cb = kwargs["callback"]
|
||||
sigmas = args[3]
|
||||
state = {"model": self.guider.model_patcher, "applied_loras": {}}
|
||||
|
||||
def step_callback(*args, **kwargs):
|
||||
apply_lora_for_step(
|
||||
self.schedules,
|
||||
args[0],
|
||||
len(sigmas),
|
||||
state,
|
||||
self.guider.model_patcher,
|
||||
self.lora_cache,
|
||||
patch=True,
|
||||
)
|
||||
if orig_cb:
|
||||
return orig_cb(*args, **kwargs)
|
||||
|
||||
kwargs["callback"] = step_callback
|
||||
apply_lora_for_step(
|
||||
self.schedules, 0, len(sigmas), state, self.guider.model_patcher, self.lora_cache, patch=True
|
||||
)
|
||||
try:
|
||||
r = self.guider.sample(*args, **kwargs)
|
||||
finally:
|
||||
unpatch_model(state["model"])
|
||||
return r
|
||||
|
||||
|
||||
class ScheduleToModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"prompt_schedule": ("PROMPT_SCHEDULE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = "promptcontrol/_legacy"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, prompt_schedule):
|
||||
return (schedule_lora_common(model, prompt_schedule),)
|
||||
|
||||
|
||||
class PCSplitSampling:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"split_sampling": (["enable", "disable"],),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = "promptcontrol/_legacy"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, split_sampling):
|
||||
model = clone_model(model)
|
||||
model.model_options["pc_split_sampling"] = split_sampling == "enable"
|
||||
return (model,)
|
||||
|
||||
|
||||
class LoRAScheduler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = "promptcontrol/_donotuse"
|
||||
FUNCTION = "apply"
|
||||
|
||||
def apply(self, model, text):
|
||||
schedules = parse_prompt_schedules(text)
|
||||
return (schedule_lora_common(model, schedules),)
|
||||
@@ -1,70 +0,0 @@
|
||||
import torch
|
||||
|
||||
|
||||
# Copied and adapted from https://github.com/bvhari/ComfyUI_PerpWeight/blob/main/clipperpweight.py
|
||||
def perp_encode(clip, tokens):
|
||||
empty_tokens = clip.tokenize("")
|
||||
sdxl_flag = "g" in tokens
|
||||
empty_cond, empty_cond_pooled = clip.encode_from_tokens(empty_tokens, return_pooled=True)
|
||||
unweighted_tokens = {}
|
||||
for k in ["l", "g"]:
|
||||
if k not in tokens:
|
||||
continue
|
||||
unweighted_tokens[k] = [[(t, 1.0) for t, _ in x] for x in tokens[k]]
|
||||
unweighted_cond, unweighted_pooled = clip.encode_from_tokens(unweighted_tokens, return_pooled=True)
|
||||
cond = torch.clone(unweighted_cond)
|
||||
|
||||
if sdxl_flag:
|
||||
for i in range(unweighted_cond.shape[0]):
|
||||
for j in range(unweighted_cond.shape[1]):
|
||||
weight_l = tokens["l"][(j // 77)][(j % 77)][1]
|
||||
if weight_l != 1.0:
|
||||
token_vector_l = unweighted_cond[i][j][:768]
|
||||
zero_vector_l = empty_cond[0][(j % 77)][:768]
|
||||
perp_l = (
|
||||
(torch.mul(zero_vector_l, token_vector_l).sum()) / (torch.norm(token_vector_l) ** 2)
|
||||
) * token_vector_l
|
||||
if weight_l > 1.0:
|
||||
cond[i][j][:768] = token_vector_l + (weight_l * perp_l)
|
||||
elif (weight_l > 0.0) and (weight_l < 1.0):
|
||||
cond[i][j][:768] = token_vector_l - ((1 - weight_l) * perp_l)
|
||||
elif weight_l < 0.0:
|
||||
cond[i][j][:768] = token_vector_l + (weight_l * perp_l)
|
||||
elif weight_l == 0.0:
|
||||
cond[i][j][:768] = empty_cond[0][(j % 77)][:768]
|
||||
|
||||
weight_g = tokens["g"][(j // 77)][(j % 77)][1]
|
||||
if weight_g != 1.0:
|
||||
token_vector_g = unweighted_cond[i][j][768:]
|
||||
zero_vector_g = empty_cond[0][(j % 77)][768:]
|
||||
perp_g = (
|
||||
(torch.mul(zero_vector_g, token_vector_g).sum()) / (torch.norm(token_vector_g) ** 2)
|
||||
) * token_vector_g
|
||||
if weight_g > 1.0:
|
||||
cond[i][j][768:] = token_vector_g + (weight_g * perp_g)
|
||||
elif (weight_g > 0.0) and (weight_g < 1.0):
|
||||
cond[i][j][768:] = token_vector_g - ((1 - weight_g) * perp_g)
|
||||
elif weight_g < 0.0:
|
||||
cond[i][j][768:] = token_vector_g + (weight_g * perp_g)
|
||||
elif weight_g == 0.0:
|
||||
cond[i][j][768:] = empty_cond[0][(j % 77)][768:]
|
||||
else:
|
||||
tokens = tokens["l"]
|
||||
for i in range(unweighted_cond.shape[0]):
|
||||
for j in range(unweighted_cond.shape[1]):
|
||||
weight = tokens[(j // 77)][(j % 77)][1]
|
||||
if weight != 1.0:
|
||||
token_vector = unweighted_cond[i][j]
|
||||
zero_vector = empty_cond[0][(j % 77)]
|
||||
perp = (
|
||||
(torch.mul(zero_vector, token_vector).sum()) / (torch.norm(token_vector) ** 2)
|
||||
) * token_vector
|
||||
if weight > 1.0:
|
||||
cond[i][j] = token_vector + (weight * perp)
|
||||
elif (weight > 0.0) and (weight < 1.0):
|
||||
cond[i][j] = token_vector - ((1 - weight) * perp)
|
||||
elif weight < 0.0:
|
||||
cond[i][j] = token_vector + (weight * perp)
|
||||
elif weight == 0.0:
|
||||
cond[i][j] = empty_cond[0][(j % 77)]
|
||||
return cond, unweighted_pooled
|
||||
@@ -1,265 +0,0 @@
|
||||
from collections import namedtuple
|
||||
from os import environ
|
||||
from math import lcm
|
||||
import time
|
||||
import logging
|
||||
import torch
|
||||
|
||||
from ..utils import lora_name_to_file, safe_float
|
||||
|
||||
import nodes
|
||||
|
||||
import comfy.model_management
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
FORCE_CPU_OFFLOAD = bool(environ.get("COMFYUI_PC_CPU_OFFLOAD"))
|
||||
|
||||
|
||||
# Minimal Modelpatcher that doesn't do anything, for LoRA loading when not
|
||||
# interested in either CLIP or unet
|
||||
class DummyModelPatcher:
|
||||
class DummyTorchModel:
|
||||
def __init__(self):
|
||||
dummyconf = {
|
||||
"num_res_blocks": [],
|
||||
"channel_mult": [],
|
||||
"transformer_depth": [],
|
||||
"transformer_depth_output": [],
|
||||
"transformer_depth_middle": 0,
|
||||
}
|
||||
self.model_config = namedtuple("DummyConfig", ["unet_config"])(dummyconf)
|
||||
|
||||
def state_dict(self):
|
||||
return {}
|
||||
|
||||
def __init__(self):
|
||||
self.model = self.DummyTorchModel()
|
||||
self.cond_stage_model = self.DummyTorchModel()
|
||||
self.weight_inplace_update = True
|
||||
self.model_options = {}
|
||||
|
||||
def add_patches(self, patches, *args, **kwargs):
|
||||
return []
|
||||
|
||||
def patch_model(self):
|
||||
pass
|
||||
|
||||
def unpatch_model(self):
|
||||
pass
|
||||
|
||||
def clone(self):
|
||||
return self
|
||||
|
||||
|
||||
DUMMY_MODEL = DummyModelPatcher()
|
||||
|
||||
|
||||
def equalize(*tensors):
|
||||
if all(t.shape[1] == tensors[0].shape[1] for t in tensors):
|
||||
return tensors
|
||||
|
||||
x = lcm(*(t.shape[1] for t in tensors))
|
||||
|
||||
return (t.repeat(1, x // t.shape[1], 1) for t in tensors)
|
||||
|
||||
|
||||
def unpatch_model(model):
|
||||
if model:
|
||||
log.info("Unpatching model")
|
||||
model.unpatch_model()
|
||||
|
||||
|
||||
def clone_model(model):
|
||||
if not model:
|
||||
return None
|
||||
model = model.clone()
|
||||
if not environ.get("PC_NO_INPLACE_UPDATE"):
|
||||
model.weight_inplace_update = True
|
||||
return model
|
||||
|
||||
|
||||
def add_patches(model, patches, weight):
|
||||
model.add_patches(patches, weight)
|
||||
|
||||
|
||||
def patch_model(model, forget=False, orig=None):
|
||||
global FORCE_CPU_OFFLOAD
|
||||
try:
|
||||
return _patch_model(model, forget, orig, FORCE_CPU_OFFLOAD)
|
||||
except comfy.model_management.OOM_EXCEPTION:
|
||||
FORCE_CPU_OFFLOAD = True
|
||||
log.error("Ran out of memory while applying LoRAs, Forcing CPU offload from now on")
|
||||
# Unpatch to restore partially applied weights
|
||||
unpatch_model(model)
|
||||
raise
|
||||
|
||||
|
||||
def _patch_model(model, forget=False, orig=None, offload_to_cpu=False):
|
||||
if not model:
|
||||
return None
|
||||
if offload_to_cpu:
|
||||
saved_offload = model.offload_device
|
||||
model.offload_device = torch.device("cpu")
|
||||
log.info(
|
||||
"Patching model, model.load_device=%s model.model.device=%s cpu_offload=%s",
|
||||
model.load_device,
|
||||
model.model.device,
|
||||
model.offload_device == torch.device("cpu"),
|
||||
)
|
||||
if orig:
|
||||
model.backup = orig.backup
|
||||
model.patch_model(device_to=model.load_device)
|
||||
if offload_to_cpu:
|
||||
model.offload_device = saved_offload
|
||||
if forget:
|
||||
model.patches = {}
|
||||
model.object_patches = {}
|
||||
return model
|
||||
|
||||
|
||||
def get_callback(model):
|
||||
return model.model_options.get("prompt_control_callback")
|
||||
|
||||
|
||||
def set_callback(model, cb):
|
||||
model.model_options["prompt_control_callback"] = cb
|
||||
|
||||
|
||||
# Hack to temporarily override printing to stdout to stop log spam
|
||||
def suppress_print(f):
|
||||
def noop(*args):
|
||||
pass
|
||||
|
||||
p = print
|
||||
__builtins__["print"] = noop
|
||||
rootlogger = logging.getLogger()
|
||||
oldlevel = rootlogger.level
|
||||
try:
|
||||
rootlogger.setLevel(logging.ERROR)
|
||||
x = f()
|
||||
except BaseException:
|
||||
__builtins__["print"] = p
|
||||
rootlogger.setLevel(oldlevel)
|
||||
raise
|
||||
__builtins__["print"] = p
|
||||
rootlogger.setLevel(oldlevel)
|
||||
return x
|
||||
|
||||
|
||||
def load_lbw():
|
||||
return nodes.NODE_CLASS_MAPPINGS.get("LoraLoaderBlockWeight //Inspire")
|
||||
|
||||
|
||||
def make_loader(filename, lbw):
|
||||
if not lbw:
|
||||
l = nodes.LoraLoader()
|
||||
|
||||
def loader(model, clip, model_weight, clip_weight, lbw):
|
||||
return suppress_print(lambda: l.load_lora(model, clip, filename, model_weight, clip_weight))
|
||||
|
||||
else:
|
||||
# This is already checked before calling make_loader
|
||||
l = load_lbw()()
|
||||
|
||||
def loader(model, clip, model_weight, clip_weight, lbw):
|
||||
spec = lbw["LBW"]
|
||||
lbw_a = safe_float(lbw.get("A"), 4.0)
|
||||
lbw_b = safe_float(lbw.get("B"), 1.0)
|
||||
m = model or DUMMY_MODEL
|
||||
c = clip or DUMMY_MODEL
|
||||
m, c, _ = suppress_print(
|
||||
lambda: l.doit(m, c, filename, model_weight, clip_weight, False, 0, lbw_a, lbw_b, "", spec)
|
||||
)
|
||||
if m is DUMMY_MODEL:
|
||||
m = None
|
||||
if c is DUMMY_MODEL:
|
||||
c = None
|
||||
return m, c
|
||||
|
||||
return loader
|
||||
|
||||
|
||||
def apply_loras_from_spec(
|
||||
loraspec, model=None, clip=None, orig_model=None, orig_clip=None, patch=False, cache=None, applied_loras=None
|
||||
):
|
||||
if applied_loras is None:
|
||||
applied_loras = {}
|
||||
actual_loraspec = {}
|
||||
additive = True
|
||||
for key in loraspec:
|
||||
if key in applied_loras and applied_loras[key] == loraspec[key]:
|
||||
continue
|
||||
if key in applied_loras and applied_loras[key] != loraspec[key]:
|
||||
additive = False
|
||||
actual_loraspec[key] = loraspec[key]
|
||||
|
||||
for key in applied_loras:
|
||||
if key not in loraspec:
|
||||
actual_loraspec = loraspec
|
||||
additive = False
|
||||
|
||||
backup_model = model
|
||||
if not additive:
|
||||
unpatch_model(model)
|
||||
# Reset clip to unpatched
|
||||
if clip:
|
||||
clip = orig_clip or clip
|
||||
|
||||
if cache is None:
|
||||
cache = {}
|
||||
if not loraspec:
|
||||
return model, clip
|
||||
|
||||
for name, params in actual_loraspec.items():
|
||||
m, c = model, clip
|
||||
w, w_clip = params["weight"], params["weight_clip"]
|
||||
if w == 0:
|
||||
m = None
|
||||
if w_clip == 0:
|
||||
c = None
|
||||
if not w and not c:
|
||||
continue
|
||||
|
||||
lbw = params.get("lbw")
|
||||
if lbw and not load_lbw():
|
||||
log.warning("LoraBlockWeight not available, ignoring LBW parameters")
|
||||
lbw = None
|
||||
|
||||
# Cache the loader instance so that it doesn't reload the LoRA from disk all the time
|
||||
cache_key = name, bool(lbw)
|
||||
loader = cache.get(cache_key)
|
||||
if not loader:
|
||||
f = lora_name_to_file(name)
|
||||
if not f:
|
||||
log.warning("Lora %s not found", name)
|
||||
continue
|
||||
log.info("Loading LoRA: %s", f)
|
||||
loader = make_loader(f, bool(lbw))
|
||||
cache[cache_key] = loader
|
||||
|
||||
m, c = loader(m, c, w, w_clip, lbw)
|
||||
model = m or model
|
||||
clip = c or clip
|
||||
if model:
|
||||
log.info("Applying LoRA: %s:%s, LBW=%s, additive=%s", name, params["weight"], bool(lbw), additive)
|
||||
if clip:
|
||||
log.info("Applying CLIP LoRA: %s:%s, LBW=%s, additive=%s", name, params["weight_clip"], bool(lbw), additive)
|
||||
|
||||
# forget patches so we don't double-patch
|
||||
model = patch_model(model, forget=True, orig=backup_model)
|
||||
return model, clip
|
||||
|
||||
|
||||
class Timer:
|
||||
def __init__(self, name):
|
||||
self.name = name
|
||||
self.start = None
|
||||
|
||||
def __enter__(self):
|
||||
self.start = time.time()
|
||||
|
||||
def __exit__(self, exc_type, exc_val, exc_tb):
|
||||
elapsed = time.time() - self.start
|
||||
if environ.get("PC_SHOW_TIMINGS"):
|
||||
log.info("Executed %s in %s seconds", self.name, elapsed)
|
||||
@@ -1,5 +1,5 @@
|
||||
import logging
|
||||
from ..parser import parse_prompt_schedules
|
||||
from .parser import parse_prompt_schedules
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
@@ -151,3 +151,15 @@ class PromptToSchedule:
|
||||
def parse(self, text, settings=None):
|
||||
schedules = parse_prompt_schedules(text)
|
||||
return (schedules,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"PCPromptToSchedule": PromptToSchedule,
|
||||
"PCScheduleSettings": PCScheduleSettings,
|
||||
"PCScheduleAddMasks": PCScheduleAddMasks,
|
||||
"PCApplySettings": PCApplySettings,
|
||||
"PCPromptFromSchedule": PCPromptFromSchedule,
|
||||
"PCFilterSchedule": FilterSchedule,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
Reference in New Issue
Block a user