Convert tests to pytest
Not 100% sure these fully work yet
This commit is contained in:
@@ -4,6 +4,9 @@ all: format check test
|
||||
check:
|
||||
ty check && ruff check
|
||||
|
||||
fix:
|
||||
ruff check --fix
|
||||
|
||||
format:
|
||||
ruff format
|
||||
|
||||
|
||||
@@ -28,6 +28,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
WEB_DIRECTORY = "web"
|
||||
|
||||
nodes = ["base", "lazy", "tools", "hooks"]
|
||||
if "PYTEST_CURRENT_TEST" in os.environ:
|
||||
nodes = []
|
||||
|
||||
for node in nodes:
|
||||
mod = importlib.import_module(f".prompt_control.nodes_{node}", package=__name__)
|
||||
|
||||
@@ -40,3 +40,6 @@ select = [
|
||||
"SIM",
|
||||
"I",
|
||||
]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
testpaths = ["tests"]
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
import logging
|
||||
|
||||
|
||||
def pytest_runtest_setup(item):
|
||||
logging.getLogger("comfyui-prompt-control").setLevel(logging.CRITICAL)
|
||||
@@ -0,0 +1,207 @@
|
||||
import numpy.testing as npt
|
||||
import pytest
|
||||
|
||||
|
||||
def run(f, *args):
|
||||
if hasattr(f, "execute"):
|
||||
return f.execute(*args)
|
||||
else:
|
||||
return getattr(f, f.FUNCTION)(*args)
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def text_encoder_clips():
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from comfy.sd import load_clip
|
||||
|
||||
clips = []
|
||||
to_test = os.environ.get("TEST_TE", "clip_l").split()
|
||||
model_dir = os.environ.get("COMFYUI_TE_DIR", ".")
|
||||
te_root = Path(model_dir).resolve()
|
||||
|
||||
if "clip_l" in to_test:
|
||||
clip_l = load_clip(
|
||||
ckpt_paths=[str(te_root / "clip_l.safetensors")], clip_type="stable_diffusion", model_options={}
|
||||
)
|
||||
clips.append(("clip_l", clip_l))
|
||||
|
||||
if "t5" in to_test:
|
||||
dual = load_clip(
|
||||
[str(te_root / "clip_l.safetensors"), str(te_root / "t5xxl_fp16.safetensors")],
|
||||
clip_type="flux",
|
||||
model_options={},
|
||||
)
|
||||
clips.append(("clip_l+t5", dual))
|
||||
|
||||
return clips
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def pc_text_encode():
|
||||
from prompt_control.nodes_base import PCTextEncode
|
||||
|
||||
return PCTextEncode()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def node_class_objs():
|
||||
import comfy_extras.nodes_mask
|
||||
import nodes
|
||||
|
||||
# Return all used node class objects
|
||||
return {
|
||||
"comfy": nodes.CLIPTextEncode(),
|
||||
"combine": nodes.ConditioningCombine(),
|
||||
"average": nodes.ConditioningAverage(),
|
||||
"concat": nodes.ConditioningConcat(),
|
||||
"zeroout": nodes.ConditioningZeroOut(),
|
||||
"strength": nodes.ConditioningSetAreaStrength(),
|
||||
"solidmask": comfy_extras.nodes_mask.SolidMask(),
|
||||
"setmask": nodes.ConditioningSetMask(),
|
||||
}
|
||||
|
||||
|
||||
def tensors_equal(t1, t2):
|
||||
npt.assert_equal(t1.detach().numpy(), t2.detach().numpy())
|
||||
|
||||
|
||||
def cond_equal(c1, c2, key=None, key_assert=None):
|
||||
assert len(c1) == len(c2)
|
||||
for i in range(len(c1)):
|
||||
a, b = c1[i], c2[i]
|
||||
if key:
|
||||
(key_assert or assert_equal)(a[1].get(key), b[1].get(key))
|
||||
else:
|
||||
tensors_equal(a[0], b[0])
|
||||
|
||||
|
||||
def assert_equal(a, b):
|
||||
assert a == b
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("text_encoder_clips", "pc_text_encode", "node_class_objs")
|
||||
class TestPCTextEncode:
|
||||
def test_basic_encode(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
combine = node_class_objs["combine"]
|
||||
average = node_class_objs["average"]
|
||||
concat = node_class_objs["concat"]
|
||||
zeroout = node_class_objs["zeroout"]
|
||||
|
||||
for _k, clip in text_encoder_clips:
|
||||
# No exceptions
|
||||
run(
|
||||
pc_text_encode,
|
||||
clip,
|
||||
"test AND test (test:1.2) BREAK test AND TE_WEIGHT(all=0) SDXL() AND AREA(,,) test CAT test",
|
||||
)
|
||||
|
||||
# Basic
|
||||
(c1,) = run(pc_text_encode, clip, "test")
|
||||
(c2,) = run(comfy, clip, "test")
|
||||
c = c2 # Used in later tests
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Quotes
|
||||
(c1,) = run(pc_text_encode, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
|
||||
(c2,) = run(comfy, clip, 'Text saying "DOG MASK AND CAT COUPLE MASK(X)"')
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Function cornercase
|
||||
(c1,) = run(pc_text_encode, clip, "test SDXL function")
|
||||
(c2,) = run(comfy, clip, "test SDXL function")
|
||||
(c3,) = run(pc_text_encode, clip, "test SDXL() function")
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Weights
|
||||
(c1,) = run(pc_text_encode, clip, "(test:1.2) (test:0.6)")
|
||||
(c2,) = run(comfy, clip, "(test:1.2) (test:0.6)")
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Concat
|
||||
(c1,) = run(pc_text_encode, clip, "test CAT test")
|
||||
(c2,) = run(concat, c, c)
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Combine
|
||||
(c1,) = run(pc_text_encode, clip, "test AND test")
|
||||
(c2,) = run(combine, c, c)
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Zero out
|
||||
(c1,) = run(pc_text_encode, clip, "test TE_WEIGHT(all=0)")
|
||||
(c2,) = run(zeroout, c)
|
||||
cond_equal(c1, c2)
|
||||
|
||||
# Average
|
||||
(c1,) = run(comfy, clip, "test1")
|
||||
(c2,) = run(comfy, clip, "test2")
|
||||
(c3,) = run(pc_text_encode, clip, "test1 AVG() test2")
|
||||
(c4,) = run(pc_text_encode, clip, "test1 AVG test2")
|
||||
(avg,) = run(average, c1, c2, 0.5)
|
||||
cond_equal(avg, c3)
|
||||
cond_equal(avg, c4)
|
||||
|
||||
@pytest.mark.xfail
|
||||
def test_failure(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c1,) = run(comfy, clip, "test SDXL function")
|
||||
(c2,) = run(pc_text_encode, clip, "test SDXL() function")
|
||||
cond_equal(c1, c2)
|
||||
|
||||
def test_weight(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
combine = node_class_objs["combine"]
|
||||
strength = node_class_objs["strength"]
|
||||
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c,) = run(comfy, clip, "test")
|
||||
(c2,) = run(strength, c, 0.5)
|
||||
# Conditioning weights
|
||||
(a,) = run(pc_text_encode, clip, "test :0.5 AND test :0.5")
|
||||
(b,) = run(combine, c2, c2)
|
||||
cond_equal(a, b)
|
||||
cond_equal(a, b, "strength")
|
||||
# Weight == 0
|
||||
(a,) = run(pc_text_encode, clip, "test :0.5 AND test :0 AND test")
|
||||
(b,) = run(combine, c2, c)
|
||||
cond_equal(a, b)
|
||||
cond_equal(a, b, "strength")
|
||||
|
||||
def test_attn_couple(self, text_encoder_clips, pc_text_encode):
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c,) = run(pc_text_encode, clip, "test COUPLE prompt1 AND test2 COUPLE prompt2")
|
||||
(c2,) = run(pc_text_encode, clip, "test COUPLE prompt1 COUPLE test2 COUPLE prompt2")
|
||||
assert len(c) == 2
|
||||
assert len(c2) == 1
|
||||
|
||||
def test_styles(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
for _k, clip in text_encoder_clips:
|
||||
(no_weights,) = run(comfy, clip, "this prompt has no weights")
|
||||
for style in ["comfy", "A1111", "comfy++", "compel", "down_weight", "perp"]:
|
||||
# no weights equal comfy
|
||||
(c,) = run(pc_text_encode, clip, "this prompt has no weights")
|
||||
cond_equal(no_weights, c)
|
||||
# does not fail when encoding weights
|
||||
for normalization in ["none", "mean", "length", "mean+length", "length+mean"]:
|
||||
run(
|
||||
pc_text_encode,
|
||||
clip,
|
||||
f"STYLE({style}, {normalization}) (this prompt) (has weights:0.9), (a:1.2) (b:1.2)",
|
||||
)
|
||||
# Just checking for exceptions
|
||||
|
||||
def test_masks(self, text_encoder_clips, pc_text_encode, node_class_objs):
|
||||
comfy = node_class_objs["comfy"]
|
||||
solidmask = node_class_objs["solidmask"]
|
||||
setmask = node_class_objs["setmask"]
|
||||
for _k, clip in text_encoder_clips:
|
||||
(c1,) = run(pc_text_encode, clip, "test MASK()")
|
||||
(c2,) = run(comfy, clip, "test")
|
||||
(c2,) = run(setmask, c2, run(solidmask, 1.0, 512, 512)[0], "default", 1.0)
|
||||
cond_equal(c1, c2)
|
||||
cond_equal(c1, c2, "mask", tensors_equal)
|
||||
@@ -0,0 +1,238 @@
|
||||
import logging
|
||||
|
||||
import pytest
|
||||
from comfy_execution.graph_utils import GraphBuilder
|
||||
|
||||
from prompt_control.nodes_lazy import (
|
||||
PCLazyLoraLoader,
|
||||
PCLazyLoraLoaderAdvanced,
|
||||
PCLazyTextEncode,
|
||||
PCLazyTextEncodeAdvanced,
|
||||
)
|
||||
|
||||
log = logging.getLogger("comfyui-prompt-control")
|
||||
|
||||
|
||||
def reset_graphbuilder_state():
|
||||
GraphBuilder.set_default_prefix("UID", 0, 0)
|
||||
|
||||
|
||||
def find_file(name):
|
||||
names = {"test": "test.safetensors", "other": "some/other.safetensors"}
|
||||
return names.get(name)
|
||||
|
||||
|
||||
def loraloader(text, adv=False, **kwargs):
|
||||
reset_graphbuilder_state()
|
||||
cls = PCLazyLoraLoader if adv else PCLazyLoraLoaderAdvanced
|
||||
model = [0, 1]
|
||||
clip = [0, 0]
|
||||
return cls().apply(unique_id="UID", model=model, clip=clip, text=text, **kwargs)
|
||||
|
||||
|
||||
def te(text, adv=False, **kwargs):
|
||||
cls = PCLazyTextEncode if adv else PCLazyTextEncodeAdvanced
|
||||
reset_graphbuilder_state()
|
||||
clip = [0, 0]
|
||||
return cls().apply(clip=clip, text=text, unique_id="UID", **kwargs)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_lora_name_to_file(monkeypatch):
|
||||
import prompt_control.utils
|
||||
|
||||
monkeypatch.setattr(prompt_control.utils, "lora_name_to_file", find_file)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def patch_torch_cuda_current_device(monkeypatch):
|
||||
import torch.cuda
|
||||
|
||||
monkeypatch.setattr(torch.cuda, "current_device", lambda: "cpu")
|
||||
|
||||
|
||||
def test_textencode_expansion():
|
||||
for p in ["test", "[test:0.2] test", "[test[test::0.5]]<lora:test:1>"]:
|
||||
r1 = te(p)
|
||||
r2 = te(p, adv=True)
|
||||
assert r1 == r2
|
||||
|
||||
|
||||
def test_textencode_lora():
|
||||
reset_graphbuilder_state()
|
||||
r = te("test<lora:test:1>")
|
||||
assert r == {
|
||||
"result": (["UID.0.0.2", 0],),
|
||||
"expand": {
|
||||
"UID.0.0.1": {"class_type": "PCTextEncode", "inputs": {"clip": [0, 0], "text": "test"}},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 1.0},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_textencode_lora_with_schedule():
|
||||
r = te("simple [test:0.1,0.5] prompt<lora:test:1>")
|
||||
assert r == {
|
||||
"result": (["UID.0.0.8", 0],),
|
||||
"expand": {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {"clip": [0, 0], "text": "simple prompt"},
|
||||
},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {"conditioning": ["UID.0.0.1", 0], "start": 0.0, "end": 0.1},
|
||||
},
|
||||
"UID.0.0.3": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {"clip": [0, 0], "text": "simple test prompt"},
|
||||
},
|
||||
"UID.0.0.4": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {"conditioning": ["UID.0.0.3", 0], "start": 0.1, "end": 0.5},
|
||||
},
|
||||
"UID.0.0.5": {
|
||||
"class_type": "PCTextEncode",
|
||||
"inputs": {"clip": [0, 0], "text": "simple prompt"},
|
||||
},
|
||||
"UID.0.0.6": {
|
||||
"class_type": "ConditioningSetTimestepRange",
|
||||
"inputs": {"conditioning": ["UID.0.0.5", 0], "start": 0.5, "end": 1.0},
|
||||
},
|
||||
"UID.0.0.7": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {"conditioning_1": ["UID.0.0.2", 0], "conditioning_2": ["UID.0.0.4", 0]},
|
||||
},
|
||||
"UID.0.0.8": {
|
||||
"class_type": "ConditioningCombine",
|
||||
"inputs": {"conditioning_1": ["UID.0.0.7", 0], "conditioning_2": ["UID.0.0.6", 0]},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_empty(monkeypatch, caplog):
|
||||
result = loraloader("prompt here <lora:nonexistent:1.0:0.5>")["expand"]
|
||||
result_adv = loraloader("prompt here <lora:nonexistent:1.0:0.5>", adv=True)["expand"]
|
||||
assert result == {}
|
||||
assert result_adv == {}
|
||||
|
||||
|
||||
def test_loraloader_duplicate_results():
|
||||
result = loraloader("<lora:test:1>")["expand"]
|
||||
result2 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>")["expand"]
|
||||
result3 = loraloader("prompt here <lora:test:1.0:0.5><lora:test:0:0.5>", adv=True)["expand"]
|
||||
assert result == result2
|
||||
assert result2 == result3
|
||||
assert result == {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": [0, 1],
|
||||
"clip": [0, 0],
|
||||
"strength_model": 1.0,
|
||||
"strength_clip": 1.0,
|
||||
"lora_name": "test.safetensors",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_multiple_loras():
|
||||
result = loraloader("<lora:test:1><lora:other:0.5>")["expand"]
|
||||
assert result == {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": [0, 1],
|
||||
"clip": [0, 0],
|
||||
"strength_model": 1.0,
|
||||
"strength_clip": 1.0,
|
||||
"lora_name": "test.safetensors",
|
||||
},
|
||||
},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": ["UID.0.0.1", 0],
|
||||
"clip": ["UID.0.0.1", 1],
|
||||
"strength_model": 0.5,
|
||||
"strength_clip": 0.5,
|
||||
"lora_name": "some/other.safetensors",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_strength_clip():
|
||||
result = loraloader("prompt here <lora:test:1.0:0.5>")["expand"]
|
||||
assert result == {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": [0, 1],
|
||||
"clip": [0, 0],
|
||||
"strength_model": 1.0,
|
||||
"strength_clip": 0.5,
|
||||
"lora_name": "test.safetensors",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_scheduled_compare():
|
||||
result = loraloader("prompt [<lora:test:0.5>:0.5]")["expand"]
|
||||
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True)["expand"]
|
||||
assert result == result2
|
||||
expected = {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "CreateHookLora",
|
||||
"inputs": {"lora_name": "test.safetensors", "strength_model": 0.5, "strength_clip": 0.5},
|
||||
},
|
||||
"UID.0.0.2": {
|
||||
"class_type": "CreateHookKeyframe",
|
||||
"inputs": {"strength_mult": 0.0, "start_percent": 0.0},
|
||||
},
|
||||
"UID.0.0.3": {
|
||||
"class_type": "CreateHookKeyframe",
|
||||
"inputs": {"start_percent": 0.5, "prev_hook_kf": ["UID.0.0.2", 0], "strength_mult": 1.0},
|
||||
},
|
||||
"UID.0.0.4": {
|
||||
"class_type": "SetHookKeyframes",
|
||||
"inputs": {"hooks": ["UID.0.0.1", 0], "hook_kf": ["UID.0.0.3", 0]},
|
||||
},
|
||||
"UID.0.0.5": {
|
||||
"class_type": "SetClipHooks",
|
||||
"inputs": {
|
||||
"clip": [0, 0],
|
||||
"hooks": ["UID.0.0.4", 0],
|
||||
"apply_to_conds": True,
|
||||
"schedule_clip": True,
|
||||
},
|
||||
},
|
||||
}
|
||||
assert result == expected
|
||||
|
||||
|
||||
def test_loraloader_adv_start():
|
||||
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", adv=True, start=0.6)["expand"]
|
||||
assert result2 == {
|
||||
"UID.0.0.1": {
|
||||
"class_type": "LoraLoader",
|
||||
"inputs": {
|
||||
"model": [0, 1],
|
||||
"clip": [0, 0],
|
||||
"strength_model": 0.5,
|
||||
"strength_clip": 0.5,
|
||||
"lora_name": "test.safetensors",
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
def test_loraloader_end_zero():
|
||||
result2 = loraloader("prompt [<lora:test:0.5>:0.5]", end=0.5)["expand"]
|
||||
assert result2 == {}
|
||||
@@ -0,0 +1,263 @@
|
||||
import pytest
|
||||
|
||||
from prompt_control.parser import expand_macros
|
||||
from prompt_control.parser import parse_prompt_schedules as parse
|
||||
|
||||
|
||||
def lora_dict(*loras):
|
||||
return {lora: {"weight": unet, "weight_clip": te} for lora, unet, te in loras}
|
||||
|
||||
|
||||
def prompt(until, text, *loras):
|
||||
return [until, {"prompt": text, "loras": lora_dict(*loras)}]
|
||||
|
||||
|
||||
def assert_prompt(p, at, until, text, *loras):
|
||||
assert p.at_step(at) == prompt(until, text, *loras)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
||||
def test_no_scheduling(step):
|
||||
p = parse("This is a (basic:0.6) (prompt) with [no scheduling] features")
|
||||
expected = prompt(1.0, "This is a (basic:0.6) (prompt) with [no scheduling] features")
|
||||
assert p.at_step(step) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
||||
def test_quote(step):
|
||||
p = parse('This is a text with a "QUOTED DEF(X=Y)"')
|
||||
expected = prompt(1.0, 'This is a text with a "QUOTED DEF(X=Y)"')
|
||||
assert p.at_step(step) == expected
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"group",
|
||||
[
|
||||
["[a:0.1]", "[:a:0.1]", "[:a::0.1,1.0]", "[:a::0.1,1.0]", "[:a::0.1]"],
|
||||
["[before:during:after:0.1]", "[before:during:after:0.1,1.0]", "[before:during:0.1]"],
|
||||
["[a:0.1,0.5]", "[[a:0.1]::0.5]", "[:a::0.1,0.5]", "[a::0.1,0.5]"],
|
||||
["[a:b:0.5]", "[a::b:0.5,0.5]"],
|
||||
["[a::0.5]", "[a:::0.5,0.5]"],
|
||||
],
|
||||
)
|
||||
def test_equivalences(group):
|
||||
objects = [parse(g) for g in group]
|
||||
first = objects[0].parsed_prompt
|
||||
for obj in objects[1:]:
|
||||
assert obj.parsed_prompt == first
|
||||
|
||||
|
||||
def test_basic():
|
||||
p = parse(
|
||||
"This is a (basic:0.6) (prompt) with (very [[simple]:(basic:0.6):0.5]:1.1) [features::0.8][ and this is ignored:1]"
|
||||
)
|
||||
assert_prompt(p, 0, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
|
||||
assert_prompt(p, 0.5, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
|
||||
assert_prompt(p, 0.7, 0.8, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) features")
|
||||
assert_prompt(p, 1.0, 1.0, "This is a (basic:0.6) (prompt) with (very (basic:0.6):1.1) ")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
||||
def test_lora(step):
|
||||
p = parse("This is a (lora:0.6) (prompt) with [no scheduling] features <lora:foo:0.5> <lora:bar:0.5:1.0>")
|
||||
expected = prompt(
|
||||
1.0, "This is a (lora:0.6) (prompt) with [no scheduling] features ", ("foo", 0.5, 0.5), ("bar", 0.5, 1.0)
|
||||
)
|
||||
assert p.at_step(step) == expected
|
||||
|
||||
|
||||
def test_scheduled_lora():
|
||||
p = parse(
|
||||
"This is a (lora:0.6) (prompt) with [scheduling] features [<lora:foo:0.5>:<lora:bar:0.5:0.2>:0.3] <lora:bar:0.5:1.0>"
|
||||
)
|
||||
assert_prompt(
|
||||
p, 0.1, 0.3, "This is a (lora:0.6) (prompt) with [scheduling] features ", ("foo", 0.5, 0.5), ("bar", 0.5, 1.0)
|
||||
)
|
||||
assert_prompt(p, 0.5, 1.0, "This is a (lora:0.6) (prompt) with [scheduling] features ", ("bar", 1.0, 1.2))
|
||||
|
||||
|
||||
def test_seq():
|
||||
p = parse("This is a sequence of [SEQ:a:0.2::0.5:c:0.8][SEQ: and x:0.8]")
|
||||
p2 = parse("This is a sequence of [[a:[c:0.5]:0.2]::0.8][ and x::0.8]")
|
||||
prompts = {
|
||||
0.2: "This is a sequence of a and x",
|
||||
0.5: "This is a sequence of and x",
|
||||
0.8: "This is a sequence of c and x",
|
||||
1.0: "This is a sequence of ",
|
||||
}
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
for k, v in prompts.items():
|
||||
assert_prompt(p, k, k, v)
|
||||
|
||||
|
||||
def test_shortcuts_scheduling():
|
||||
p = parse("A schedule [a:0.1,0.7] b")
|
||||
p2 = parse("A schedule [[a:0.1]::0.7] b")
|
||||
p3 = parse("A schedule [a:b:0.5,0.8]")
|
||||
p4 = parse("A schedule [[a:0.5]:b:0.8]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
assert p3.parsed_prompt == p4.parsed_prompt
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"step,until,text",
|
||||
[
|
||||
(0, 0.1, "test excluded test"),
|
||||
(0.2, 0.4, "test test"),
|
||||
(0.45, 1.0, "test excluded2 test"),
|
||||
],
|
||||
)
|
||||
def test_range_1(step, until, text):
|
||||
p = parse("test [excluded::excluded2:0.1,0.4] test")
|
||||
assert_prompt(p, step, until, text)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"step,until,text",
|
||||
[
|
||||
(0, 0.1, "test test"),
|
||||
(0.25, 0.3, "test included test"),
|
||||
(0.15, 0.2, "test excluded test"),
|
||||
(0.55, 0.6, "test test"),
|
||||
(0.95, 1.0, "test excluded2 test"),
|
||||
],
|
||||
)
|
||||
def test_range_2(step, until, text):
|
||||
p = parse("test [[:included::0.2,0.8]|[excluded::excluded2:0.4,0.9]:0.1] test")
|
||||
assert_prompt(p, step, until, text)
|
||||
|
||||
|
||||
def test_nested():
|
||||
p = parse(
|
||||
"This [prompt is [SEQ:[crazy:weird:0.2] stuff:0.5:<lora:cool:1>:0.7:nesting:1.0]:completely ignored with tags:HR]"
|
||||
)
|
||||
prompts = {
|
||||
0.2: (0.2, "This prompt is crazy stuff"),
|
||||
0.3: (0.5, "This prompt is weird stuff"),
|
||||
0.5: (0.5, "This prompt is weird stuff"),
|
||||
0.8: (1.0, "This prompt is nesting"),
|
||||
}
|
||||
for k in prompts:
|
||||
exp = [prompts[k][0], {"prompt": prompts[k][1], "loras": {}}]
|
||||
assert p.at_step(k) == exp
|
||||
|
||||
assert_prompt(p, 0.6, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
|
||||
assert_prompt(p, 0.7, 0.7, "This prompt is ", ("cool", 1.0, 1.0))
|
||||
p2 = p.with_filters(filters="hr, xyz")
|
||||
assert p2.at_step(0) == p2.at_step(1)
|
||||
|
||||
|
||||
def test_def():
|
||||
p = parse("DEF(X=0.5) [a:b:X] DEF(test = [c:X]) test test")
|
||||
cases = [
|
||||
(0.2, 0.5, "a "),
|
||||
(0.6, 1.0, "b c c"),
|
||||
]
|
||||
for k, until, text in cases:
|
||||
assert_prompt(p, k, until, text)
|
||||
|
||||
p = parse("DEF(X=[($1):($1:$2):$2])X(test;0.7)")
|
||||
p2 = parse("[(test):(test:0.7):0.7]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
p = parse("DEF(X=[($1):($1:$2):$2])DEF(Y=X(test;$1))Y(0.7) Y(0.5)")
|
||||
p2 = parse("[(test):(test:0.7):0.7] [(test):(test:0.5):0.5]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
p = expand_macros("DEF(X(a;b)=$1 $2 $3 d)X(A) X(A;B;C)")
|
||||
assert p == "A b $3 d A B C d"
|
||||
|
||||
p = expand_macros("DEF(MACRO()=[empty:$1:$2])MACRO MACRO(;) MACRO(;0.5) MACRO(a;0.5)")
|
||||
assert p == "[empty::$2] [empty::] [empty::0.5] [empty:a:0.5]"
|
||||
|
||||
p = expand_macros("DEF(X=$1)DEF(Y()=$1)[X Y][X() Y()][X(1) Y(1)]")
|
||||
assert p == "[$1 ][ ][1 1]"
|
||||
|
||||
p = parse("DEF(test(1)=prompt $1)DEF(test2((a); (test))=[$1:$2:0.5])test test2")
|
||||
p2 = parse("prompt 1 [(a):(prompt 1):0.5]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
with pytest.raises(ValueError) as c:
|
||||
expand_macros("DEF(X=recurse Y) DEF(Y=recurse X) X")
|
||||
assert "Unable to resolve DEFs" in str(c.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"text, cases",
|
||||
[
|
||||
(r"[a:\:a:0.5] :\[a:b:0.5]", [(0, 0.5, r"a :\[a:b:0.5]"), (0.55, 1, r":a :\[a:b:0.5]")]),
|
||||
(r"[embedding\:a:embedding\:b:0.1,0.5]", [(0.15, 0.5, r"embedding:a"), (0.55, 1, r"embedding:b")]),
|
||||
(
|
||||
r"[embedding\:a:embedding\:b:embedding\:c:0.1,0.5]",
|
||||
[(0.0, 0.1, r"embedding:a"), (0.15, 0.5, r"embedding:b"), (0.55, 1, r"embedding:c")],
|
||||
),
|
||||
(r"[a\:b\\:c:0.5]", [(0.0, 0.5, "a:b\\"), (0.55, 1, r"c")]),
|
||||
(r"[a:\#b:0.5]", [(0.0, 0.5, "a"), (0.55, 1, "#b")]),
|
||||
],
|
||||
)
|
||||
def test_escapes(text, cases):
|
||||
p = parse(text)
|
||||
for step, until, val in cases:
|
||||
assert_prompt(p, step, until, val)
|
||||
|
||||
|
||||
def test_comments():
|
||||
p = parse("this is a # comment")
|
||||
assert_prompt(p, 0, 1.0, "this is a ")
|
||||
p = parse("this is a [comment#:scheduled:0.6]")
|
||||
assert_prompt(p, 0, 1.0, "this is a [comment")
|
||||
p = parse(r"this is a [comment\#:scheduled:0.6]")
|
||||
assert_prompt(p, 0, 0.6, "this is a comment#")
|
||||
assert_prompt(p, 0.65, 1.0, "this is a scheduled")
|
||||
p = parse("#this is a comment\nthis is a prompt")
|
||||
assert_prompt(p, 0, 1.0, "\nthis is a prompt")
|
||||
|
||||
|
||||
def test_misc():
|
||||
p = parse("[[a:c:0.5]:0.7]")
|
||||
p2 = parse("[:[a:c:0.5]:0.7]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
p = parse("test [[a:[b<lora:test:0.5>:0.6]:0.5]:HR]")
|
||||
p2 = parse("test [:[a:[:b<lora:test:0.5>:0.6]:0.5]:HR]")
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
pf = p.with_filters(filters="hr")
|
||||
assert pf.parsed_prompt == p2.with_filters(filters="hr").parsed_prompt
|
||||
assert_prompt(pf, 0, 0.5, "test a")
|
||||
assert_prompt(pf, 0.55, 0.6, "test ")
|
||||
assert_prompt(pf, 0.8, 1.0, "test b", ("test", 0.5, 0.5))
|
||||
|
||||
p = parse("[:[<lora:test:1>:c:0.5]:0.3]")
|
||||
assert_prompt(p, 0, 0.3, "")
|
||||
assert_prompt(p, 0.4, 0.5, "", ("test", 1.0, 1.0))
|
||||
assert_prompt(p, 1.0, 1.0, "c")
|
||||
|
||||
p = parse("an [<emb:foo>:<emb:bar>:0.5]")
|
||||
prompts = {
|
||||
0.2: (0.5, "an embedding:foo"),
|
||||
0.8: (1.0, "an embedding:bar"),
|
||||
}
|
||||
for k, (until, val) in prompts.items():
|
||||
assert_prompt(p, k, until, val)
|
||||
|
||||
|
||||
def test_alternating():
|
||||
p = parse("[cat|dog|tiger]")
|
||||
p2 = parse("[cat|dog|tiger:0.1]")
|
||||
p3 = parse("[cat|[dog|wolf]|tiger]")
|
||||
p4 = parse("[cat|[dog:wolf<lora:canine:1>:0.5]:0.2]")
|
||||
|
||||
assert p.parsed_prompt == p2.parsed_prompt
|
||||
|
||||
catdogtigers = ["cat", "wolf", "tiger", "cat", "dog", "tiger", "cat", "wolf", "tiger", "cat"]
|
||||
for i, x in enumerate(catdogtigers):
|
||||
step = round((i * 0.1) + 0.1, 2)
|
||||
assert_prompt(p3, step, step, x)
|
||||
|
||||
for i, (text, *_loras) in enumerate(
|
||||
[(["cat"],), (["dog"],), (["cat"],), (["wolf", ("canine", 1.0, 1.0)],), (["cat"],)]
|
||||
):
|
||||
step = round((i * 0.2) + 0.2, 2)
|
||||
assert_prompt(p4, step, step, *text)
|
||||
assert_prompt(p4, 0.7, 0.8, "wolf", ("canine", 1.0, 1.0))
|
||||
Reference in New Issue
Block a user