Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8aec0d8f46 | ||
|
|
013586904b | ||
|
|
85ff8466ad | ||
|
|
97d17d2bfd | ||
|
|
c0e671b2de | ||
|
|
7ff55c6717 | ||
|
|
524738f21d | ||
|
|
5c7d507e91 |
@@ -1 +1,2 @@
|
||||
__pycache__
|
||||
.pyre
|
||||
|
||||
@@ -197,6 +197,7 @@ class AdvancedEncoder:
|
||||
def add_encoder(cls, name, fn):
|
||||
cls.STYLES[name] = fn
|
||||
|
||||
@classmethod
|
||||
def add_normalization_op(cls, name, fn):
|
||||
cls.NORMALIZATION_OPS[name] = fn
|
||||
|
||||
@@ -296,7 +297,7 @@ class AdvancedEncoder:
|
||||
w_mix = np.diff([0] + w.tolist())
|
||||
w_mix = torch.tensor(w_mix, dtype=embs.dtype, device=embs.device).reshape((-1, 1, 1))
|
||||
|
||||
weighted_emb = (w_mix * embs).sum(axis=0, keepdim=True)
|
||||
weighted_emb = (w_mix * embs).sum(dim=0, keepdim=True)
|
||||
pooled = pooled_base
|
||||
if pooled is not None and self.max_length:
|
||||
pooled = weighted_emb[0, self.max_length - 1 : self.max_length, :]
|
||||
@@ -328,12 +329,13 @@ class AdvancedEncoder:
|
||||
masks = torch.cat(masks)
|
||||
|
||||
embs = base_emb.expand(embs.shape) - embs
|
||||
pooled = None
|
||||
if pooled_base is not None and self.max_length:
|
||||
pooled = embs[0, self.max_length - 1 : self.max_length, :]
|
||||
pooled_start = pooled_base.expand(len(ws), -1)
|
||||
ws = torch.tensor(ws).reshape(-1, 1).expand(pooled_start.shape)
|
||||
pooled = (pooled - pooled_start) * (ws - 1)
|
||||
pooled = pooled.mean(axis=0, keepdim=True)
|
||||
pooled = pooled.mean(dim=0, keepdim=True)
|
||||
pooled = pooled_base + pooled
|
||||
|
||||
if embs.shape[0] != masks.shape[0]:
|
||||
|
||||
@@ -71,14 +71,13 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
}
|
||||
}
|
||||
self.has_negpip = False
|
||||
|
||||
# calculate later. All clones must refer to the same kv dict
|
||||
self.kv = {"k": None, "v": None}
|
||||
self.kv = {}
|
||||
|
||||
def initialize_regions(self, base_cond, conds, fill):
|
||||
self.num_conds = len(conds) + 1
|
||||
self.base_strength = base_cond[1].get("strength", 1.0)
|
||||
self.strengths = [cond[1].get("strength", 1.0) for cond in conds]
|
||||
self.strengths: list[float] = [cond[1].get("strength", 1.0) for cond in conds]
|
||||
self.conds: list[torch.Tensor] = [base_cond[0]] + [cond[0] for cond in conds]
|
||||
base_mask = base_cond[1].get("mask", None)
|
||||
masks = [cond[1].get("mask") * cond[1].get("mask_strength") for cond in conds]
|
||||
@@ -215,6 +214,7 @@ class AttentionCoupleHook(TransformerOptionsHook):
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
assert self.num_conds is not None, "this is a bug"
|
||||
cond_or_uncond_couple.extend(itertools.repeat(self.COND, self.num_conds))
|
||||
|
||||
q = torch.cat(qs, dim=0)
|
||||
|
||||
@@ -1,47 +0,0 @@
|
||||
import comfy_execution.caching
|
||||
from comfy_execution.graph_utils import is_link
|
||||
import nodes
|
||||
from os import environ
|
||||
|
||||
import logging
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
include_unique_id_in_input = comfy_execution.caching.include_unique_id_in_input
|
||||
|
||||
|
||||
def promptcontrol_get_immediate_node_signature(self, dynprompt, node_id, ancestor_order_mapping):
|
||||
if not dynprompt.has_node(node_id):
|
||||
# This node doesn't exist -- we can't cache it.
|
||||
return [float("NaN")]
|
||||
node = dynprompt.get_node(node_id)
|
||||
class_type = node["class_type"]
|
||||
class_def = nodes.NODE_CLASS_MAPPINGS[class_type]
|
||||
inputs = node["inputs"]
|
||||
if hasattr(class_def, "CACHE_KEY"):
|
||||
inputs = getattr(class_def, "CACHE_KEY")(inputs)
|
||||
signature = [class_type, self.is_changed_cache.get(node_id)]
|
||||
if (
|
||||
self.include_node_id_in_input()
|
||||
or (hasattr(class_def, "NOT_IDEMPOTENT") and class_def.NOT_IDEMPOTENT)
|
||||
or include_unique_id_in_input(class_type)
|
||||
):
|
||||
signature.append(node_id)
|
||||
for key in sorted(inputs.keys()):
|
||||
if is_link(inputs[key]):
|
||||
(ancestor_id, ancestor_socket) = inputs[key]
|
||||
ancestor_index = ancestor_order_mapping[ancestor_id]
|
||||
signature.append((key, ("ANCESTOR", ancestor_index, ancestor_socket)))
|
||||
else:
|
||||
signature.append((key, inputs[key]))
|
||||
return signature
|
||||
|
||||
|
||||
def init():
|
||||
if environ.get("PROMPTCONTROL_ENABLE_CACHE_HACK") != "1":
|
||||
return
|
||||
log.warning("Enabling Prompt Control cache hack")
|
||||
comfy_execution.caching.CacheKeySetInputSignature.get_immediate_node_signature = (
|
||||
promptcontrol_get_immediate_node_signature
|
||||
)
|
||||
@@ -215,7 +215,7 @@ def encode_regions(clip_regions, encode, tokenizer):
|
||||
region_emb *= region_masking
|
||||
|
||||
region_embeddings.append(region_emb)
|
||||
region_embeddings = torch.stack(region_embeddings).sum(axis=0)
|
||||
region_embeddings = torch.stack(region_embeddings).sum(dim=0)
|
||||
|
||||
embeddings_final_mask = torch.tensor(
|
||||
global_region_mask, dtype=base_embedding_full.dtype, device=base_embedding_full.device
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
from .prompts import encode_prompt
|
||||
|
||||
@@ -40,7 +41,7 @@ class PCTextEncode:
|
||||
DESCRIPTION = "Encodes a prompt with extra goodies from Prompt Control. This node does *not* support scheduling"
|
||||
|
||||
def apply(self, clip, text):
|
||||
return PCTextEncodeWithRange.apply(self, clip, text, 0.0, 1.0)
|
||||
return PCTextEncodeWithRange().apply(clip, text, 0.0, 1.0)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"PCTextEncode": PCTextEncode, "PCTextEncodeWithRange": PCTextEncodeWithRange}
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
|
||||
import comfy.hooks
|
||||
@@ -38,7 +39,6 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
all_hooks = []
|
||||
|
||||
def create_hook(loraspec, start_pct, end_pct, non_scheduled):
|
||||
nonlocal lora_cache
|
||||
hooks = []
|
||||
hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||
for path, info in loras.items():
|
||||
@@ -53,7 +53,8 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
lora_cache[path], strength_model=info["weight"], strength_clip=info["weight_clip"]
|
||||
)
|
||||
# Set hook_ref so that identical hooks compare equal
|
||||
new_hook.hooks[0].hook_ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
||||
ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
|
||||
new_hook.hooks[0].hook_ref = ref # pyright: ignore[reportAttributeAccessIssue]
|
||||
hooks.append(new_hook)
|
||||
if start_pct > 0.0:
|
||||
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
|
||||
@@ -74,8 +75,6 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
all_hooks.append(hook)
|
||||
start_pct = end_pct
|
||||
|
||||
del lora_cache
|
||||
|
||||
all_hooks = [x for x in all_hooks if x]
|
||||
|
||||
if all_hooks:
|
||||
@@ -85,7 +84,7 @@ def lora_hooks_from_schedule(schedules, non_scheduled):
|
||||
|
||||
class PCAttentionCoupleBatchNegative(ComfyNodeABC):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> InputTypeDict:
|
||||
def INPUT_TYPES(s) -> InputTypeDict:
|
||||
return {
|
||||
"required": {
|
||||
"positive": (IO.CONDITIONING, {}),
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
from __future__ import annotations
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules
|
||||
from comfy_execution.graph_utils import GraphBuilder, is_link
|
||||
@@ -167,7 +169,7 @@ class PCLazyLoraLoaderAdvanced:
|
||||
"hidden": {"unique_id": "UNIQUE_ID"},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL", "CLIP", "HOOKS")
|
||||
RETURN_TYPES: tuple[str, ...] = ("MODEL", "CLIP", "HOOKS")
|
||||
OUTPUT_TOOLTIPS = ("Returns a model and clip with LoRAs scheduled",)
|
||||
CATEGORY = "promptcontrol"
|
||||
FUNCTION = "apply"
|
||||
@@ -214,8 +216,7 @@ def build_scheduled_prompts(graph, schedules, clip):
|
||||
classname = "PCTextEncode"
|
||||
paramname = "text"
|
||||
if classnames:
|
||||
classname = classnames[0][0]
|
||||
paramname = classnames[0][1]
|
||||
classname, paramname = classnames[0].args
|
||||
node = graph.node(classname)
|
||||
node.set_input("clip", clip)
|
||||
node.set_input(paramname, p)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# pyright: reportSelfClsParameterName=false
|
||||
import logging
|
||||
from .parser import parse_prompt_schedules, expand_macros
|
||||
from .nodes_lazy import NODE_CLASS_MAPPINGS as LAZY_NODES
|
||||
|
||||
+20
-11
@@ -1,4 +1,5 @@
|
||||
# vim: sw=4 ts=4
|
||||
from __future__ import annotations
|
||||
import lark
|
||||
import logging
|
||||
from math import ceil
|
||||
@@ -100,7 +101,15 @@ class CutTransform(lark.Transformer):
|
||||
def cut(self, args):
|
||||
prompt, cutout, weight, strict_mask, start_from_masked, mask_token = args
|
||||
|
||||
return ("".join(flatten(prompt)), "".join(flatten(cutout)), weight, strict_mask, start_from_masked, mask_token)
|
||||
# prompts and cutouts are always sequences of str
|
||||
return (
|
||||
"".join(flatten(prompt)), # pyright: ignore
|
||||
"".join(flatten(cutout)), # pyright: ignore
|
||||
weight,
|
||||
strict_mask,
|
||||
start_from_masked,
|
||||
mask_token,
|
||||
) # pyright: ignore
|
||||
|
||||
def start(self, args):
|
||||
prompt = []
|
||||
@@ -301,8 +310,7 @@ def at_step(step, filters, tree):
|
||||
return name, params, lbw
|
||||
|
||||
def __default__(self, data, children, meta):
|
||||
for child in children:
|
||||
yield child
|
||||
return children
|
||||
|
||||
return AtStep().transform(tree)
|
||||
|
||||
@@ -427,7 +435,9 @@ def expand_macros(text):
|
||||
prevres = text
|
||||
replacements = []
|
||||
for d in defs:
|
||||
r = d.split("=", 1)
|
||||
if not d.args:
|
||||
continue
|
||||
r = d.args[0].split("=", 1)
|
||||
search = parse_search(r[0].strip())
|
||||
if not search or len(r) != 2:
|
||||
log.warning("Ignoring invalid DEF(%s)", d)
|
||||
@@ -452,15 +462,14 @@ def expand_macros(text):
|
||||
|
||||
def substitute_defcall(text, search, replace):
|
||||
name, default_args = search
|
||||
text, defns = get_function(
|
||||
text, name, defaults=None, placeholder=f"DEFNCALL{name}", require_args=False, return_dict=True
|
||||
)
|
||||
text, defns = get_function(text, name, defaults=None, placeholder=f"DEFNCALL{name}", require_args=False)
|
||||
for i, d in enumerate(defns):
|
||||
ph = d["placeholder"]
|
||||
parameters = d["args"]
|
||||
ph = d.placeholder
|
||||
assert ph is not None, "This is a bug"
|
||||
parameters = d.args
|
||||
paramvals = []
|
||||
if parameters is not None:
|
||||
paramvals = [x.strip() for x in parameters.split(";")]
|
||||
if parameters:
|
||||
paramvals = [x.strip() for x in parameters[0].split(";")]
|
||||
r = replace
|
||||
for i, v in enumerate(paramvals):
|
||||
r = re.sub(rf"\${i+1}\b", v, r)
|
||||
|
||||
+43
-29
@@ -1,3 +1,4 @@
|
||||
from __future__ import annotations
|
||||
import logging
|
||||
import re
|
||||
import torch
|
||||
@@ -6,7 +7,17 @@ from functools import partial
|
||||
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
|
||||
from nodes import ConditioningAverage
|
||||
|
||||
from .utils import safe_float, get_function, split_by_function, parse_floats, smarter_split, call_node, split_quotable
|
||||
from .utils import (
|
||||
safe_float,
|
||||
get_function,
|
||||
split_by_function,
|
||||
parse_floats,
|
||||
smarter_split,
|
||||
call_node,
|
||||
split_quotable,
|
||||
FunctionSpec,
|
||||
ComfyConditioning,
|
||||
)
|
||||
from .adv_encode import advanced_encode_from_tokens
|
||||
from .cutoff import process_cuts
|
||||
from .parser import parse_cuts
|
||||
@@ -26,7 +37,7 @@ def get_sdxl(text, defaults):
|
||||
text, sdxl = get_function(text, "SDXL", ["none", "none", "none"])
|
||||
if not sdxl:
|
||||
return text, {}
|
||||
args = sdxl[0]
|
||||
args = sdxl[0].args
|
||||
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+")
|
||||
@@ -47,7 +58,7 @@ def get_clipweights(text, existing_spec=None):
|
||||
text, spec = get_function(text, "TE_WEIGHT", defaults=None)
|
||||
if not spec:
|
||||
return existing_spec or {}, text
|
||||
args = spec[0].strip()
|
||||
args = spec[0].args[0].strip()
|
||||
res = {}
|
||||
for arg in args.split(","):
|
||||
try:
|
||||
@@ -63,7 +74,7 @@ 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, normalization = styles[0].args
|
||||
style = style.strip()
|
||||
normalization = normalization.strip()
|
||||
if style.replace("old+", "") not in AVAILABLE_STYLES:
|
||||
@@ -78,8 +89,9 @@ def get_style(text, default_style="comfy", default_normalization="none"):
|
||||
return style, normalization, text
|
||||
|
||||
|
||||
def shuffle_chunk(shuffle, c):
|
||||
func, shuffle = shuffle
|
||||
def shuffle_chunk(func_spec: FunctionSpec, c: str) -> str:
|
||||
func = func_spec.name
|
||||
shuffle = func_spec.args
|
||||
shuffle_count = int(safe_float(shuffle[0], 0))
|
||||
_, separator, joiner = shuffle
|
||||
if separator == "default":
|
||||
@@ -129,11 +141,11 @@ def fix_word_ids(tokens):
|
||||
|
||||
|
||||
def tokenize_chunks(clip, text, need_word_ids, can_break):
|
||||
chunks = split_quotable(text, r"\bBREAK\b")
|
||||
chunks = list(split_quotable(text, r"\bBREAK\b"))
|
||||
token_chunks = []
|
||||
shuffled_chunks = []
|
||||
for c in chunks:
|
||||
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"], return_func_name=True)
|
||||
c, shuffles = get_function(c.strip(), "(SHIFT|SHUFFLE)", ["0", "default", "default"])
|
||||
r = c
|
||||
for s in shuffles:
|
||||
r = shuffle_chunk(s, r)
|
||||
@@ -169,9 +181,10 @@ def tokenize(clip, text, can_break, empty_tokens):
|
||||
per_te_prompts = {}
|
||||
if l_prompts:
|
||||
log.warning("Note: CLIP_L is deprecated. Use TE(l=prompt) instead")
|
||||
per_te_prompts["l"] = l_prompts
|
||||
per_te_prompts["l"] = [x.args for x in l_prompts]
|
||||
|
||||
for prompt in te_prompts:
|
||||
prompt = prompt.args[0]
|
||||
if prompt.strip() == "help":
|
||||
log.info("Encoders available for TE: %s", ", ".join(tokens.keys()))
|
||||
continue
|
||||
@@ -212,7 +225,7 @@ def encode_prompt_segment(
|
||||
default_style="comfy",
|
||||
default_normalization="none",
|
||||
clip_weights=None,
|
||||
) -> list[tuple[torch.Tensor, dict[str]]]:
|
||||
) -> list[ComfyConditioning]:
|
||||
style, normalization, text = get_style(text, default_style, default_normalization)
|
||||
clip_weights, text = get_clipweights(text, clip_weights)
|
||||
text, cuts = parse_cuts(text)
|
||||
@@ -234,17 +247,16 @@ def encode_prompt_segment(
|
||||
|
||||
text, averages = split_by_function(text, "AVG", ["0.5"], require_args=False)
|
||||
prompts_to_avg = []
|
||||
for avg in averages:
|
||||
w = safe_float(avg["args"][0], 0.5)
|
||||
for chunk, avg in averages:
|
||||
w = safe_float(avg.args[0], 0.5)
|
||||
prompts_to_avg.append((text, w))
|
||||
text = avg["text"]
|
||||
text = chunk
|
||||
prompts_to_avg.append((text, 1.0))
|
||||
|
||||
conds_to_avg = []
|
||||
for prompt, weight in prompts_to_avg:
|
||||
conds_to_cat = []
|
||||
chunks = split_quotable(prompt, r"\bCAT\b")
|
||||
for c in chunks:
|
||||
for c in split_quotable(prompt, r"\bCAT\b"):
|
||||
tokens = tokenize(clip, c, can_break, empty)
|
||||
conds_to_cat.append(clip.encode_from_tokens_scheduled(tokens, add_dict=settings))
|
||||
|
||||
@@ -366,7 +378,7 @@ def get_area(text):
|
||||
if not areas:
|
||||
return text, None
|
||||
|
||||
args = areas[0]
|
||||
args = areas[0].args
|
||||
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)
|
||||
@@ -393,7 +405,7 @@ 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]
|
||||
w, h = sizes[0].args
|
||||
return text, (int(w), int(h))
|
||||
|
||||
|
||||
@@ -446,14 +458,14 @@ def get_mask(text, size, input_masks):
|
||||
mask = None
|
||||
totalweight = 1.0
|
||||
if maskw:
|
||||
totalweight = safe_float(maskw[0][0], 1.0)
|
||||
totalweight = safe_float(maskw[0].args[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)
|
||||
weight = safe_float(m.args[2], 1.0)
|
||||
op = m.args[3]
|
||||
nextmask = make_mask(m.args, size, weight)
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
nextmask = feather(feathers[i].args, nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
log.info("MaskComposite op=%s", op)
|
||||
@@ -461,7 +473,8 @@ def get_mask(text, size, input_masks):
|
||||
else:
|
||||
mask = nextmask
|
||||
|
||||
for idx, w, op in imasks:
|
||||
for im in imasks:
|
||||
idx, w, op = im.args
|
||||
idx = int(safe_float(idx, 0.0))
|
||||
w = safe_float(w, 1.0)
|
||||
if input_masks is None:
|
||||
@@ -475,7 +488,7 @@ def get_mask(text, size, input_masks):
|
||||
continue
|
||||
nextmask = input_masks[idx] * w
|
||||
if i < len(feathers):
|
||||
nextmask = feather(feathers[i], nextmask)
|
||||
nextmask = feather(feathers[i].args, nextmask)
|
||||
i += 1
|
||||
if mask is not None:
|
||||
mask = call_node(MaskComposite, mask, nextmask, 0, 0, op)[0]
|
||||
@@ -484,7 +497,7 @@ def get_mask(text, size, input_masks):
|
||||
|
||||
# apply leftover FEATHER() specs to the whole
|
||||
for f in feathers[i:]:
|
||||
mask = feather(f, mask)
|
||||
mask = feather(f.args, mask)
|
||||
|
||||
return text, mask, totalweight
|
||||
|
||||
@@ -499,14 +512,15 @@ def get_noise(text):
|
||||
return text, None, None
|
||||
w = 0
|
||||
# Only take seed from first noise spec, for simplicity
|
||||
seed = safe_float(noises[0][1], "none")
|
||||
seed = noises[0].args[0].strip()
|
||||
if seed == "none":
|
||||
gen = None
|
||||
else:
|
||||
seed = safe_float(seed, 0)
|
||||
gen = torch.Generator()
|
||||
gen.manual_seed(int(seed))
|
||||
for n in noises:
|
||||
w += safe_float(n[0], 0.0)
|
||||
w += safe_float(n.args[0], 0.0)
|
||||
return text, max(min(w, 1.0), 0.0), gen
|
||||
|
||||
|
||||
@@ -567,7 +581,7 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
|
||||
style, normalization, text = get_style(text)
|
||||
text, mask_size = get_mask_size(text, defaults)
|
||||
|
||||
prompts = split_quotable(text, r"\bAND\b")
|
||||
prompts = list(split_quotable(text, r"\bAND\b"))
|
||||
|
||||
p, sdxl_opts = get_sdxl(prompts[0], defaults)
|
||||
prompts[0] = p
|
||||
@@ -591,7 +605,7 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks):
|
||||
for prompt in prompts:
|
||||
base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False)
|
||||
|
||||
prompts = [base_prompt] + [couple_mask(p["args"]) + p["text"] for p in attn_couple_prompts]
|
||||
prompts = [base_prompt] + [couple_mask(f.args) + chunk for (chunk, f) in attn_couple_prompts]
|
||||
encoded = []
|
||||
for p in prompts:
|
||||
p, settings = process_settings(p, defaults, masks, mask_size, sdxl_opts)
|
||||
|
||||
@@ -20,11 +20,9 @@ def run(f, *args):
|
||||
return getattr(f, f.FUNCTION)(*args)
|
||||
|
||||
|
||||
@mock.patch("torch.cuda.current_device", lambda: "cpu")
|
||||
class TestEncode(unittest.TestCase):
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
global clips
|
||||
print("Loading ComfyUI")
|
||||
from comfy.sd import load_clip
|
||||
from pathlib import Path
|
||||
|
||||
+61
-47
@@ -1,15 +1,34 @@
|
||||
from __future__ import annotations
|
||||
from pathlib import Path
|
||||
import re
|
||||
import logging
|
||||
import copy
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, TypeAlias, Iterator, TypeVar, TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
import torch # flakes8: noqa
|
||||
|
||||
FunctionArgs: TypeAlias = list[str]
|
||||
ComfyConditioning: TypeAlias = tuple["torch.Tensor", dict[str, Any]]
|
||||
|
||||
|
||||
@dataclass
|
||||
class FunctionSpec:
|
||||
name: str
|
||||
args: FunctionArgs
|
||||
position: int
|
||||
placeholder: str | None
|
||||
|
||||
|
||||
# Allow testing
|
||||
try:
|
||||
from folder_paths import get_filename_list
|
||||
except ImportError:
|
||||
|
||||
def get_filename_list(x):
|
||||
raise NotImplementedError("How did you get here?")
|
||||
def get_filename_list(folder_name) -> list[str]:
|
||||
return []
|
||||
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
@@ -64,10 +83,11 @@ def find_nonscheduled_loras(consolidated_schedule):
|
||||
return {k: v for (k, v) in candidate_loras.items() if k not in to_remove}
|
||||
|
||||
|
||||
def smarter_split(separator, string):
|
||||
def smarter_split(separator: str, string: str) -> list[str]:
|
||||
"""Does not break () when splitting"""
|
||||
splits = []
|
||||
prev = 0
|
||||
idx = 0
|
||||
stack = 0
|
||||
escape = False
|
||||
for idx, x in enumerate(string):
|
||||
@@ -84,7 +104,7 @@ def smarter_split(separator, string):
|
||||
return splits
|
||||
|
||||
|
||||
def find_closing_paren(text, start):
|
||||
def find_closing_paren(text: str, start: int) -> int:
|
||||
stack = 1
|
||||
for i, char in enumerate(text[start:]):
|
||||
if char == ")":
|
||||
@@ -96,7 +116,9 @@ def find_closing_paren(text, start):
|
||||
return -1
|
||||
|
||||
|
||||
def find_function_spans(text, func, require_args, defaults):
|
||||
def find_function_spans(
|
||||
text: str, func: str, require_args: bool, defaults: FunctionArgs | None
|
||||
) -> Iterator[tuple[int, int, str, FunctionArgs]]:
|
||||
if require_args:
|
||||
rex = re.compile(rf"\b{func}\(", re.MULTILINE)
|
||||
else:
|
||||
@@ -113,20 +135,21 @@ def find_function_spans(text, func, require_args, defaults):
|
||||
if text[at_paren:after_first_paren] == "(":
|
||||
end = find_closing_paren(text, after_first_paren)
|
||||
if end < 0:
|
||||
print("no closing paren:", text)
|
||||
continue
|
||||
args = parse_strings(text[after_first_paren:end], defaults)
|
||||
end += 1
|
||||
else:
|
||||
end = at_paren
|
||||
args = defaults
|
||||
args = defaults or []
|
||||
yield idx + start, idx + end, funcname, args
|
||||
idx = idx + end
|
||||
text = text[end:]
|
||||
match = rex.search(text)
|
||||
|
||||
|
||||
def get_function(text, func, defaults, return_func_name=False, placeholder="", return_dict=False, require_args=True):
|
||||
def get_function(
|
||||
text: str, func: str, defaults: list[str] | None, placeholder: str = "", require_args: bool = True
|
||||
) -> tuple[str, list[FunctionSpec]]:
|
||||
spans = [x.span() for x in re.finditer(r'".+?"', text)]
|
||||
instances = []
|
||||
count = 0
|
||||
@@ -138,24 +161,8 @@ def get_function(text, func, defaults, return_func_name=False, placeholder="", r
|
||||
continue
|
||||
if placeholder:
|
||||
ph = f"\0{placeholder}{count}\0"
|
||||
if return_dict:
|
||||
instances.append(
|
||||
{
|
||||
"name": funcname,
|
||||
"args": args,
|
||||
"position": start,
|
||||
"placeholder": ph,
|
||||
}
|
||||
)
|
||||
elif return_func_name:
|
||||
instances.append((funcname, args))
|
||||
else:
|
||||
instances.append(args)
|
||||
|
||||
if placeholder:
|
||||
chunks.append(text[current:start] + f"\0{placeholder}{count}\0")
|
||||
else:
|
||||
chunks.append(text[current:start])
|
||||
instances.append(FunctionSpec(funcname, args, start, ph))
|
||||
chunks.append(text[current:start] + (ph or ""))
|
||||
current = end
|
||||
count += 1
|
||||
chunks.append(text[current:])
|
||||
@@ -163,60 +170,67 @@ def get_function(text, func, defaults, return_func_name=False, placeholder="", r
|
||||
return text, instances
|
||||
|
||||
|
||||
def spans_include(spans, s, e):
|
||||
def spans_include(spans: list[tuple[int, int]], s: int, e: int) -> bool:
|
||||
return any((s > a and e < b) for a, b in spans)
|
||||
|
||||
|
||||
def split_quotable(text, regexp):
|
||||
res = []
|
||||
def split_quotable(text: str, regexp: str) -> Iterator[str]:
|
||||
start_from = 0
|
||||
spans = [x.span() for x in re.finditer(r'".+?"', text)]
|
||||
for x in re.finditer(regexp, text):
|
||||
s, e = x.span()
|
||||
if not spans_include(spans, s, e):
|
||||
res.append(text[start_from:s].strip())
|
||||
yield text[start_from:s].strip()
|
||||
start_from = e
|
||||
res.append(text[start_from:].strip())
|
||||
return res
|
||||
yield text[start_from:].strip()
|
||||
|
||||
|
||||
def split_by_function(text, func, defaults=None, require_args=True):
|
||||
def split_by_function(
|
||||
text: str, func: str, defaults: list[str] | None = None, require_args: bool = True
|
||||
) -> tuple[str, list[tuple[str, FunctionSpec]]]:
|
||||
"""
|
||||
Splits a string by function calls, returning the text preceding the first call and a list of dictionaries with a "text" key with the prompt before the next split or until hthe end of the text.
|
||||
Splits a string by function calls, returning the leftover text along with a list of functions with their associated text chunk.
|
||||
"""
|
||||
text, functions = get_function(text, func, defaults, return_dict=True, require_args=require_args)
|
||||
text, functions = get_function(text, func, defaults, require_args=require_args)
|
||||
chunks = []
|
||||
prev = 0
|
||||
for f in functions:
|
||||
chunks.append(text[prev : f["position"]])
|
||||
prev = f["position"]
|
||||
chunks.append(text[prev : f.position])
|
||||
prev = f.position
|
||||
chunks.append(text[prev:])
|
||||
r = []
|
||||
for i, f in enumerate(functions):
|
||||
f["text"] = chunks[i + 1]
|
||||
return chunks[0], functions
|
||||
r.append((chunks[i + 1], f))
|
||||
return chunks[0], r
|
||||
|
||||
|
||||
def parse_args(strings, arg_spec, strip=True):
|
||||
T = TypeVar("T")
|
||||
|
||||
|
||||
def parse_args(strings: list[str], arg_spec: list[tuple[Any, T]], strip: bool = True) -> list[T]:
|
||||
args = [s[1] for s in arg_spec]
|
||||
for i, spec in list(enumerate(arg_spec))[: len(strings)]:
|
||||
try:
|
||||
if strip:
|
||||
strings[i] = strings[i].strip()
|
||||
args[i] = spec[0](strings[i])
|
||||
f = spec[0]
|
||||
args[i] = f(strings[i])
|
||||
except ValueError:
|
||||
pass
|
||||
return args
|
||||
|
||||
|
||||
def parse_floats(string, defaults, split_re=","):
|
||||
def parse_floats(string: str, defaults: list[float], split_re: str = ",") -> list[float]:
|
||||
spec = [(float, d) for d in defaults]
|
||||
return parse_args(re.split(split_re, string.strip()), spec)
|
||||
|
||||
|
||||
def parse_strings(string, defaults, split_re=r"(?<!\\),", replace=(r"\,", ",")):
|
||||
def parse_strings(
|
||||
string: str, defaults: FunctionArgs | None, split_re: str = r"(?<!\\),", replace: tuple[str, str] = (r"\,", ",")
|
||||
) -> FunctionArgs:
|
||||
if defaults is None:
|
||||
return string
|
||||
spec = [(lambda x: x, d) for d in defaults]
|
||||
return [string]
|
||||
spec = [(str, d) for d in defaults]
|
||||
splits = re.split(split_re, string)
|
||||
if replace:
|
||||
f, t = replace
|
||||
@@ -224,7 +238,7 @@ def parse_strings(string, defaults, split_re=r"(?<!\\),", replace=(r"\,", ",")):
|
||||
return parse_args(splits, spec, strip=False)
|
||||
|
||||
|
||||
def safe_float(f, default):
|
||||
def safe_float(f: Any, default: float) -> float:
|
||||
if f is None:
|
||||
return default
|
||||
try:
|
||||
@@ -233,7 +247,7 @@ def safe_float(f, default):
|
||||
return default
|
||||
|
||||
|
||||
def lora_name_to_file(name):
|
||||
def lora_name_to_file(name: str) -> str | None:
|
||||
filenames = get_filename_list("loras")
|
||||
# Return exact matches as is
|
||||
if name in filenames:
|
||||
|
||||
@@ -13,3 +13,7 @@ Repository = "https://github.com/asagi4/comfyui-prompt-control"
|
||||
PublisherId = "asagi4"
|
||||
DisplayName = "ComfyUI Prompt Control"
|
||||
Icon = ""
|
||||
|
||||
[tool.pyright]
|
||||
extraPaths = ["../../"]
|
||||
exclude = ["prompt_control/*test*"]
|
||||
|
||||
Reference in New Issue
Block a user