Compare commits

..
15 Commits
Author SHA1 Message Date
asagi4 67d41fb1b3 v1.2.0 2024-12-03 14:24:10 +02:00
asagi4 71e340939b Document new nodes 2024-12-03 14:21:54 +02:00
asagi4 751af8cabb Initial nodes using the new hooks mechanism recently merged 2024-12-03 14:05:54 +02:00
asagi4 7e9ca60dfd Pad with an empty prompt
See #64
2024-11-27 00:02:42 +02:00
asagi4 8b76376e56 Update example.json 2024-09-23 18:57:48 +03:00
asagi4 4bbf3a895f Release 1.1.2 2024-08-22 21:45:14 +03:00
asagi4 2930f03d6c Make sure that LoRAs are loaded to the correct device 2024-08-22 21:44:49 +03:00
asagi4 42acef7298 Make flux work 2024-08-15 10:15:35 +03:00
asagi4 7ae4882155 New release 2024-08-01 13:43:58 +03:00
asagi4 4c331b55fa Merge pull request #56 from ComfyNodePRs/licence-update
Update PyProject Toml - License
2024-08-01 13:41:39 +03:00
snomiao 4e68c07970 chore(licence-update): Update PyProject Toml - License 2024-07-31 17:02:38 +00:00
asagi4 6505c64c6e Don't use a hack to detect the sampler type
See #55
2024-07-22 12:47:47 +03:00
asagi4 ec069716e8 Only adjust word indexes when a "restart" actually happens 2024-07-20 18:38:58 +03:00
asagi4 ef9dd40985 Fix word indexes when BREAK is used.
This should fix #51
2024-07-20 18:18:20 +03:00
asagi4 b8c83c2eb0 Typo, fixes #50 2024-06-20 16:48:12 +03:00
9 changed files with 2316 additions and 1436 deletions
+36 -18
View File
@@ -13,7 +13,7 @@ Things you can control via the prompt:
- SDXL parameters
- Other miscellaneous things
[This example workflow](workflows/example.json?raw=1) implements a two-pass workflow illustrating most scheduling features.
[This example workflow](workflows/example.json?raw=1) implements a two-pass workflow illustrating most scheduling features. (Note: the example still uses deprecated nodes; update pending).
The tools in this repository combine well with the macro and wildcard functionality in [comfyui-utility-nodes](https://github.com/asagi4/comfyui-utility-nodes)
@@ -32,6 +32,7 @@ Then restart ComfyUI afterwards.
I try to avoid behavioural changes that break old prompts, but they may happen occasionally.
- 2024-12-03 ComfyUI merged support for model/conditioning hooks. There are two new nodes, `PCEncodeSchedule` and `PCLoraHooksFromSchedule` that can be used in combination with the hook nodes. Some functionality is still missing from them, but going forward, these nodes will be the only nodes supported; **I will not spend significant time fixing bugs in the old monkeypatched nodes anymore.**
- 2024-02-02 The node will now automatically enable offloading LoRA backup weights to the CPU if you run out of memory during LoRA operations, even when `--highvram` is specified. This change persists until ComfyUI is restarted.
- 2024-01-14 Multiple `CLIP_L` instances are now joined with a space separator instead of concatenated.
- 2024-01-09 AITemplate support dropped. I don't recommend or test AITemplate anymore. Use Stable-Fast instead (see below for info)
@@ -57,7 +58,7 @@ a [large::0.1] [cat|dog:0.05] [<lora:somelora:0.5:0.6>::0.5]
[in a park:in space:0.4]
```
You can also use `a [b:c:0.3,0.7]` as a shortcut. The prompt be `a` until 0.3, `a b` until 0.7, and then `c`. `[a:0.1,0.4]` is equivalent to `[a::0.1,0.4]`
You can also use `a [b:c:0.3,0.7]` as a shortcut. The prompt be `a` until 0.3, `a b` until 0.7, and then `a c`. `[a:0.1,0.4]` is equivalent to `[a::0.1,0.4]`
## LoRA loading
@@ -101,6 +102,8 @@ with `tags` `x,z` would result in the prompt `a blue cat running in space`
## Prompt interpolation
Note: Not currently supported by `PCEncodeSchedule`
`a red [INT:dog:cat:0.2,0.8:0.05]` will attempt to interpolate the tensors for `a red dog` and `a red cat` between the specified range in as many steps of 0.05 as will fit.
@@ -215,6 +218,10 @@ gives you a mask that is a combination of 1, 2 and 3, where 1 and 3 are feathere
The order of the `FEATHER` and `MASK` calls doesn't matter; you can have `FEATHER` before `MASK` or even interleave them.
# Schedulable LoRAs
Note: Use `PCLoraHooksFromSchedule`. It will work better.
## Old nodes
The `ScheduleToModel` node patches a model so that when sampling, it'll switch LoRAs between steps. You can apply the LoRA's effect separately to CLIP conditioning and the unet (model).
Swapping LoRAs often can be quite slow without the `--highvram` switch because ComfyUI will shuffle things between the CPU and GPU. When things stay on the GPU, it's quite fast.
@@ -225,6 +232,8 @@ You can also set the `PC_RETRY_ON_OOM` environment variable to any non-empty val
## LoRA Block Weight
Note: Not supported by `PCEncodeSchedule` yet
If you have [ComfyUI Inspire Pack](https://github.com/ltdrdata/ComfyUI-Inspire-Pack) installed, you can use its Lora Block Weight syntax, for example:
```
@@ -235,6 +244,8 @@ The syntax is the same as in the `ImpactWildcard` node, documented [here](https:
# Other integrations
## Advanced CLIP encoding
Note: `perp` is not supported by `PCEncodeSchedule`
You can use the syntax `STYLE(weight_interpretation, normalization)` in a prompt to affect how prompts are interpreted.
Without any extra nodes, only `perp` is available, which does the same as [ComfyUI_PerpWeight](https://github.com/bvhari/ComfyUI_PerpWeight) extension.
@@ -256,6 +267,8 @@ For things (ie. the code imports) to work, the nodes must be cloned in a directo
## Cutoff node integration
Note: Not supported by `PCEncodeSchedule` yet.
If you have [ComfyUI Cutoff](https://github.com/BlenderNeko/ComfyUI_Cutoff) cloned into your `custom_nodes`, you can use the `CUT` keyword to use cutoff functionality
The syntax is
@@ -265,12 +278,16 @@ a group of animals, [CUT:white cat:white], [CUT:brown dog:brown:0.5:1.0:1.0:_]
the parameters in the `CUT` section are `region_text:target_text:weight;strict_mask:start_from_masked:padding_token` of which only the first two are required.
If `strict_mask`, `start_from_masked` or `padding_token` are specified in more than one section, the last one takes effect for the whole prompt
## Stable-Fast
The prompt control node works well with [ComfyUI_stable_fast](https://github.com/gameltb/ComfyUI_stable_fast). However, you should apply `ScheduleToModel` **after** applying `Apply StableFast Unet` to prevent constant recompilations.
# Nodes
## PCLoraHooksFromSchedule
Creates a ComfyUI `HOOKS` object from a prompt schedule. Can be attached to a CLIP model to perform encoding and LoRA switching
## PCEncodeSchedule
Encodes all prompts in a schedule. Pass in a `CLIP` object with hooks attached for LoRA scheduling, then use the resulting `CONDITIONING` normally
## PromptToSchedule
Parses a schedule from a text prompt. A schedule is essentially an array of `(valid_until, prompt)` pairs that the other nodes can use.
@@ -283,17 +300,6 @@ Always returns at least the last prompt in the schedule if everything would othe
`start=0, end=0` returns the prompt at the start and `start=1.0, end=1.0` returns the prompt at the end.
## ScheduleToCond
Produces a combined conditioning for the appropriate timesteps. From a schedule. Also applies LoRAs to the CLIP model according to the schedule.
## ScheduleToModel
Produces a model that'll cause the sampler to reapply LoRAs at specific steps according to the schedule.
This depends on a callback handled by a monkeypatch of the ComfyUI sampler function, so it might not work with custom samplers, but it shouldn't interfere with them either.
## PCSplitSampling
Causes sampling to be split into multiple sampler calls instead of relying on timesteps for scheduling. This makes the schedules more accurate, but seems to cause weird behaviour with SDE samplers. (Upstream bug?)
## PCScheduleSettings
Returns an object representing **default values** for the `SDXL` function and allows configuring `MASK_SIZE` outside the prompt. You need to apply them to a schedule with `PCApplySettings`. Note that for the SDXL settings to apply, you still need to have `SDXL()` in the prompt.
@@ -311,7 +317,19 @@ LoRAs are *not* included in the text prompt, though they are logged.
Attaches custom masks to a `PROMPT_SCHEDULE` that can then be used in a prompt.
## PromptControlSimple
## ScheduleToCond (deprecated)
Produces a combined conditioning for the appropriate timesteps. From a schedule. Also applies LoRAs to the CLIP model according to the schedule.
## ScheduleToModel (deprecated)
Produces a model that'll cause the sampler to reapply LoRAs at specific steps according to the schedule.
This depends on a callback handled by a monkeypatch of the ComfyUI sampler function, so it might not work with custom samplers, but it shouldn't interfere with them either.
## PCSplitSampling (deprecated)
Causes sampling to be split into multiple sampler calls instead of relying on timesteps for scheduling. This makes the schedules more accurate, but seems to cause weird behaviour with SDE samplers. (Upstream bug?)
## PromptControlSimple (deprecated)
This node exists purely for convenience. It's a combination of `PromptToSchedule`, `ScheduleToCond`, `ScheduleToModel` and `FilterSchedule` such that it provides as output a model, positive conds and negative conds, both with and without any specified filters applied.
This makes it handy for quick one- or two-pass workflows.
+3
View File
@@ -13,6 +13,7 @@ from .prompt_control.node_other import (
PCPromptFromSchedule,
)
from .prompt_control.node_aio import PromptControlSimple
from .prompt_control.node_hooks import PCLoraHooksFromSchedule, PCEncodeSchedule
log = logging.getLogger("comfyui-prompt-control")
log.propagate = False
@@ -43,4 +44,6 @@ NODE_CLASS_MAPPINGS = {
"ScheduleToModel": ScheduleToModel,
"EditableCLIPEncode": EditableCLIPEncode,
"LoRAScheduler": LoRAScheduler,
"PCLoraHooksFromSchedule": PCLoraHooksFromSchedule,
"PCEncodeSchedule": PCEncodeSchedule,
}
+5 -5
View File
@@ -18,7 +18,7 @@ def hijack(obj, attr, replacement):
setattr(replacement, "pc_hijack_done", True)
def hijack_sampler(module, function):
def hijack_sampler(module, function, is_custom):
mod = sys.modules[module]
orig_sampler = getattr(mod, function)
if has_hijack(orig_sampler):
@@ -36,7 +36,7 @@ def hijack_sampler(module, function):
if cb:
try:
try:
r = cb(orig_sampler, *args, **kwargs)
r = cb(orig_sampler, is_custom, *args, **kwargs)
except comfy.model_management.OOM_EXCEPTION:
if not os.environ.get("PC_RETRY_ON_OOM"):
raise
@@ -45,7 +45,7 @@ def hijack_sampler(module, function):
BrownianTreeNoiseSampler.pc_reset(False)
gc.collect()
comfy.model_management.soft_empty_cache()
r = cb(orig_sampler, *args, **kwargs)
r = cb(orig_sampler, is_custom, *args, **kwargs)
except Exception:
log.error("Exception occurred during callback, unpatching model.")
unpatch_model(model)
@@ -155,6 +155,6 @@ def hijack_browniannoisesampler(module, cls):
def do_hijack():
hijack_browniannoisesampler("comfy.k_diffusion.sampling", "BrownianTreeNoiseSampler")
hijack_sampler("comfy.sample", "sample")
hijack_sampler("comfy.sample", "sample_custom")
hijack_sampler("comfy.sample", "sample", False)
hijack_sampler("comfy.sample", "sample_custom", True)
hijack_ksampler("comfy.samplers", "KSampler")
+26 -8
View File
@@ -286,6 +286,22 @@ def shuffle_chunk(shuffle, c):
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 = []
@@ -296,6 +312,7 @@ def encode_prompt(clip, text, default_style="comfy", default_normalization="none
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
@@ -311,28 +328,29 @@ def encode_prompt(clip, text, default_style="comfy", default_normalization="none
t = vector_sculptor_tokens(clip, c, method, norm, w)
else:
# Tokenizer returns padded results
t = clip.tokenize(c, return_word_ids=len(regions) > 0 or (have_advanced_encode and style != "perp"))
t = clip.tokenize(c, return_word_ids=need_word_ids)
token_chunks.append(t)
tokens = token_chunks[0]
for c in token_chunks[1:]:
for key in tokens:
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=len(regions) > 0 or (have_advanced_encode and style != "perp")
)["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(text_l, return_word_ids=len(regions) > 0 or (have_advanced_encode and style != "perp"))
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)
@@ -341,7 +359,7 @@ def encode_prompt(clip, text, default_style="comfy", default_normalization="none
log.warning("Normalization is not supported with perp style weighting. Ignored '%s'", normalization)
return perp_encode(clip, tokens)
if have_advanced_encode and not sculpts:
if "t5xxl" not in tokens and have_advanced_encode and not sculpts:
if "g" in tokens:
embs_l = None
embs_g = None
+508
View File
@@ -0,0 +1,508 @@
import logging
import re
import torch
from .utils import safe_float, get_function, parse_floats, lora_name_to_file
from comfy_extras.nodes_mask import FeatherMask, MaskComposite
import comfy.utils
import comfy.hooks
import folder_paths
log = logging.getLogger("comfyui-prompt-control")
AVAILABLE_STYLES = ["comfy"]
AVAILABLE_NORMALIZATIONS = ["none"]
have_advanced_encode = False
try:
import custom_nodes.ComfyUI_ADV_CLIP_emb.adv_encode as adv_encode
have_advanced_encode = True
AVAILABLE_STYLES.extend(["A1111", "compel", "comfy++", "down_weight"])
AVAILABLE_NORMALIZATIONS.extend(["mean", "length", "length+mean"])
except ImportError:
pass
class PCLoraHooksFromSchedule:
@classmethod
def INPUT_TYPES(s):
return {
"required": {"prompt_schedule": ("PROMPT_SCHEDULE",)},
}
RETURN_TYPES = ("HOOKS",)
OUTPUT_TOOLTIPS = ("set of hooks created from the prompt schedule",)
CATEGORY = "promptcontrol/_unstable"
FUNCTION = "apply"
def apply(self, prompt_schedule):
return (lora_hooks_from_schedule(prompt_schedule),)
class PCEncodeSchedule:
@classmethod
def INPUT_TYPES(s):
return {
"required": {"clip": ("CLIP",), "prompt_schedule": ("PROMPT_SCHEDULE",)},
}
RETURN_TYPES = ("CONDITIONING",)
CATEGORY = "promptcontrol/_unstable"
FUNCTION = "apply"
def apply(self, clip, prompt_schedule):
return (encode_schedule(clip, prompt_schedule),)
SHUFFLE_GEN = torch.Generator(device="cpu")
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 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, settings, default_style="comfy", default_normalization="none"
) -> list[tuple[torch.Tensor, dict[str]]]:
style, normalization, text = get_style(text, default_style, default_normalization)
# 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 = have_advanced_encode or style == "comfy" and normalization == "none"
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
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)
newclip = clip
if have_advanced_encode:
newclip = clip.clone()
if hasattr(clip.patcher.model, "clip_g"):
newclip.patcher.add_object_patch(
"clip_g.encode_token_weights",
encoder_patch(style, normalization, clip.patcher.get_model_object("clip_g.encode_token_weights")),
)
if hasattr(clip.patcher.model, "clip_l"):
newclip.patcher.add_object_patch(
"clip_l.encode_token_weights",
encoder_patch(style, normalization, clip.patcher.get_model_object("clip_l.encode_token_weights")),
)
return newclip.encode_from_tokens_scheduled(tokens, add_dict=settings)
def encoder_patch(style, normalization, orig_fn):
if not have_advanced_encode or style == "comfy" and normalization == "none":
return orig_fn
else:
log.debug("Encoding with style=%s, normalization=%s", style, normalization)
return lambda t: adv_encode.advanced_encode_from_tokens(
t, normalization, style, orig_fn, return_pooled=True, apply_to_pooled=False
)
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, start_pct, end_pct, 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)
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 = []
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)
settings = {"prompt": prompt}
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
settings["start_percent"] = start_pct
settings["end_percent"] = end_pct
x = encode_prompt(clip, prompt, settings, style, normalization)
conds.extend(x)
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 lora_hooks_from_schedule(schedules):
start_pct = 0.0
lora_cache = {}
all_hooks = []
prev_loras = {}
def create_hook(loraspec, start_pct, end_pct):
nonlocal lora_cache
hooks = []
hook_kf = comfy.hooks.HookKeyframeGroup()
for lora, info in loras.items():
path = lora_name_to_file(lora)
if not path:
continue
if path not in lora_cache:
lora_cache[path] = comfy.utils.load_torch_file(
folder_paths.get_full_path("loras", path), safe_load=True
)
new_hook = comfy.hooks.create_hook_lora(
lora_cache[path], strength_model=info["weight"], strength_clip=info["weight_clip"]
)
# Set hook_ref so that identical hooks compare equal
new_hook.hooks[0].hook_ref = f"pc-{path}-{info['weight']}-{info['weight_clip']}"
hooks.append(new_hook)
if start_pct > 0.0:
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=0.0)
hook_kf.add(kf)
kf = comfy.hooks.HookKeyframe(strength=1.0, start_percent=start_pct)
hook_kf.add(kf)
if end_pct < 1.0:
kf = comfy.hooks.HookKeyframe(strength=0.0, start_percent=end_pct)
hook_kf.add(kf)
hooks = comfy.hooks.HookGroup.combine_all_hooks(hooks)
if hooks:
hooks.set_keyframes_on_hooks(hook_kf=hook_kf)
return hooks
consolidated = []
prev_loras = {}
for end_pct, c in reversed(list(schedules)):
loras = c["loras"]
if loras != prev_loras:
consolidated.append((end_pct, loras))
prev_loras = loras
consolidated = reversed(consolidated)
for end_pct, loras in consolidated:
log.info("Creating LoRA hook from %s to %s: %s", start_pct, end_pct, loras)
hook = create_hook(loras, start_pct, end_pct)
all_hooks.append(hook)
start_pct = end_pct
del lora_cache
all_hooks = [x for x in all_hooks if x is not None]
if all_hooks:
hooks = comfy.hooks.HookGroup.combine_all_hooks(all_hooks)
return hooks
def encode_schedule(clip, schedules):
start_pct = 0.0
conds = []
for end_pct, c in schedules:
if start_pct < end_pct:
prompt = c["prompt"]
cond = do_encode(clip, prompt, start_pct, end_pct, schedules.defaults, schedules.masks)
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
+3 -4
View File
@@ -1,6 +1,5 @@
import logging
import torch
import inspect
from .utils import unpatch_model, clone_model, set_callback, apply_loras_from_spec
from .parser import parse_prompt_schedules
@@ -37,17 +36,17 @@ def schedule_lora_common(model, schedules, lora_cache=None):
if lora_cache is None:
lora_cache = {}
def sampler_cb(orig_sampler, *args, **kwargs):
def sampler_cb(orig_sampler, is_custom, *args, **kwargs):
split_sampling = args[0].model_options.get("pc_split_sampling")
state = {}
# For custom samplers, sigmas is not a keyword argument. Do the check this way to fall back to old behaviour if other hijacks exist.
if "sigmas" in inspect.getfullargspec(orig_sampler).args:
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
+7 -2
View File
@@ -173,10 +173,15 @@ def _patch_model(model, forget=False, orig=None, offload_to_cpu=False):
if offload_to_cpu:
saved_offload = model.offload_device
model.offload_device = torch.device("cpu")
log.info("Patching model, cpu_offload=%s", 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()
model.patch_model(device_to=model.load_device)
if offload_to_cpu:
model.offload_device = saved_offload
if forget:
+2 -2
View File
@@ -1,8 +1,8 @@
[project]
name = "comfyui-prompt-control"
description = "Nodes for convenient prompt editing, making many common operations prompt-controllable"
version = "1.1.0"
license = "LICENSE"
version = "1.2.0"
license = { file = "LICENSE" }
# some lark versions older than 1.1.9 apparently have a bug that breaks things, see https://github.com/asagi4/comfyui-prompt-control/issues/35
dependencies = ["lark >= 1.1.9"]
+1726 -1397
View File
File diff suppressed because it is too large Load Diff