Compare commits

...
8 Commits
Author SHA1 Message Date
asagi4 8aec0d8f46 I can't run tests if this exists... 2025-12-18 23:51:43 +02:00
asagi4 013586904b Appease typing 2025-12-18 23:51:43 +02:00
asagi4 85ff8466ad Appease typing 2025-12-18 23:51:43 +02:00
asagi4 97d17d2bfd Minor refactor to appease typing 2025-12-18 23:51:43 +02:00
asagi4 c0e671b2de More typing 2025-12-18 23:51:43 +02:00
asagi4 7ff55c6717 Make split_quotable an iterator 2025-12-16 18:08:35 +02:00
asagi4 524738f21d Add some more typing 2025-12-16 18:01:50 +02:00
asagi4 5c7d507e91 Refactor get_function to makes its use consistent
Add some typing, just for fun
2025-12-16 17:47:10 +02:00
14 changed files with 148 additions and 151 deletions
+1
View File
@@ -1 +1,2 @@
__pycache__
.pyre
+4 -2
View File
@@ -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]:
+3 -3
View File
@@ -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)
-47
View File
@@ -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
)
+1 -1
View File
@@ -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
+2 -1
View File
@@ -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}
+4 -5
View File
@@ -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, {}),
+4 -3
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)
-2
View File
@@ -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
View File
@@ -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:
+4
View File
@@ -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*"]