Add some SDXL support, should fix #12

This commit is contained in:
asagi4
2023-12-31 21:47:23 +02:00
parent 48810e6a15
commit c5ee13ee90
3 changed files with 67 additions and 8 deletions
+11
View File
@@ -75,6 +75,15 @@ a [black:blue:X] [cat:dog:Y] [walking:running:Z] in space
```
with `tags` `x,z` would result in the prompt `a blue cat running in space`
## SDXL
You can use the function `SDXL(width height, target_width target_height, crop_w crop_h)` to set SDXL prompt parameters. `SDXL()` is equivalent to `SDXL(1024 1024, 1024 1024, 0 0)`
To set the `clip_l` prompt, as with `CLIPTextEncodeSDXL`, use the function `CLIP_L(prompt text goes here)`. multiple instances of `CLIP_L` are concatenated, and `BREAK` isn't supported in it. It has no effect on SD 1.5. The rest of the prompt becomes the `clip_g` prompt.
if there is no `CLIP_L`, the prompts will work as with `CLIPTextEncode`.
## Other syntax:
- `<emb:xyz>` is alternative syntax for `embedding:xyz` to work around a syntax conflict with `[embedding:xyz:0.5]` which is parsed as a schedule that switches from `embedding` to `xyz`.
@@ -88,6 +97,8 @@ cat :1 AND dog :2
```
The weight defaults to 1 and are normalized so that `a:2 AND b:2` is equal to `a AND b`. `AND` is processed after schedule parsing, so you can change the weight mid-prompt: `cat:[1:2:0.5] AND dog`
if there is `COMFYAND()` in the prompt, the behaviour of `AND` will change to work like `ConditioningCombine`, but in practice this seems to be just slower while producing the same output.
## Functions
There are some "functions" that can be included in a prompt to do various things.
+54 -8
View File
@@ -159,6 +159,25 @@ class EditableCLIPEncode:
return (control_to_clip_common(self, clip, parsed),)
def get_sdxl(text):
text, sdxl = get_function(text, "SDXL", ["", "1024 1024", "1024 1024", "0 0"])
if not sdxl:
return text, {}
args = sdxl[0]
w, h = parse_floats(args[0], [1024, 1024], split_re="\s+")
tw, th = parse_floats(args[1], [1024, 1024], split_re="\s+")
cropw, croph = parse_floats(args[2], [0, 0], split_re="\s+")
opts = {
"width": int(w),
"height": int(h),
"target_width": int(tw),
"target_height": int(tw),
"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:
@@ -216,6 +235,8 @@ def encode_regions(clip, tokens, regions, weight_interpretation="comfy", token_n
def encode_prompt(clip, text, default_style="comfy", default_normalization="none"):
style, normalization, text = get_style(text, default_style, default_normalization)
text, regions = parse_cuts(text)
# 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 = []
for c in chunks:
@@ -225,12 +246,23 @@ def encode_prompt(clip, text, default_style="comfy", default_normalization="none
token_chunks.append(t)
tokens = token_chunks[0]
for c in token_chunks[1:]:
if isinstance(tokens, list):
tokens.extend(c)
else:
# dict, SDXL
for key in tokens:
tokens[key].extend(c[key])
for key in tokens:
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"]
if "g" 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"))
while len(tokens["l"]) < len(tokens["g"]):
tokens["l"] += empty["l"]
while len(tokens["l"]) > len(tokens["g"]):
tokens["g"] += empty["g"]
if len(regions) > 0:
return encode_regions(clip, tokens, regions, style, normalization)
@@ -385,8 +417,16 @@ def do_encode(clip, text):
# First style modifier applies to ANDed prompts too unless overridden
style, normalization, text = get_style(text)
text, mask_size = get_mask_size(text)
# Don't sum ANDs if this is in prompt
alt_method = "COMFYAND()" in text
text = text.replace("COMFYAND()", "")
prompts = [p.strip() for p in re.split(r"\bAND\b", text)]
p, sdxl_opts = get_sdxl(prompts[0])
prompts[0] = p
def weight(t):
opts = {}
m = re.search(r":(-?\d\.?\d*)(![A-Za-z]+)?$", t)
@@ -410,11 +450,17 @@ def do_encode(clip, text):
if not w:
continue
prompt, area = get_area(prompt)
prompt, local_sdxl_opts = get_sdxl(p)
cond, pooled = encode_prompt(clip, prompt, style, normalization)
cond = apply_noise(cond, noise_w, generator)
pooled = apply_noise(pooled, noise_w, generator)
settings = {"prompt": prompt}
if alt_method:
settings["strength"] = w
prompt, local_sdxl_opts = get_sdxl(p)
settings.update(sdxl_opts)
settings.update(local_sdxl_opts)
if area:
settings["area"] = area[0]
settings["strength"] = area[1]
@@ -423,7 +469,7 @@ def do_encode(clip, text):
settings["mask"] = mask
settings["mask_strength"] = mask_weight
if mask is not None or area:
if mask is not None or area or alt_method or local_sdxl_opts:
if pooled is not None:
settings["pooled_output"] = pooled
conds.append([cond, settings])
@@ -435,7 +481,7 @@ def do_encode(clip, text):
pooleds = [r[1] for r in res if r[1] is not None]
if len(res) > 0:
opts = {}
opts = sdxl_opts
if pooleds:
opts["pooled_output"] = sum(equalize(*pooleds))
sumcond = sum(equalize(*sumconds))
+2
View File
@@ -42,6 +42,8 @@ def parse_floats(string, defaults, split_re=","):
def parse_strings(string, defaults, split_re=","):
if defaults is None:
return string
spec = [(lambda x: x, d) for d in defaults]
return parse_args(re.split(split_re, string.strip()), spec)