Compare commits

..
7 Commits
Author SHA1 Message Date
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
5 changed files with 36 additions and 19 deletions
+1 -1
View File
@@ -57,7 +57,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
+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")
+25 -7
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(text_l, 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)
+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
+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.1.1"
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"]