Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4bbf3a895f | ||
|
|
2930f03d6c | ||
|
|
42acef7298 | ||
|
|
7ae4882155 | ||
|
|
4c331b55fa | ||
|
|
4e68c07970 | ||
|
|
6505c64c6e | ||
|
|
ec069716e8 | ||
|
|
ef9dd40985 | ||
|
|
b8c83c2eb0 |
@@ -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
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -173,10 +173,10 @@ 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
@@ -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.2"
|
||||
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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user