367 lines
12 KiB
Python
367 lines
12 KiB
Python
import os
|
|
|
|
import pytest
|
|
|
|
|
|
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 prompts_match(a, b):
|
|
return list(a) == list(b)
|
|
|
|
|
|
def assert_prompt(p, at, until, text, *loras):
|
|
assert prompts_match(p.at_step(at), prompt(until, text, *loras))
|
|
|
|
|
|
parsers_to_test = os.environ.get("PC_PARSERS_TO_TEST", "new").split()
|
|
|
|
params = []
|
|
if "old" in parsers_to_test:
|
|
from prompt_control.parser_lark import parse_prompt_schedules as old_parse # noqa
|
|
|
|
params.append(old_parse)
|
|
|
|
if "new" in parsers_to_test:
|
|
from prompt_control.parser_parsy import parse_prompt_schedules as new_parse # noqa
|
|
|
|
params.append(new_parse)
|
|
|
|
|
|
@pytest.fixture(scope="module", autouse=True, params=params)
|
|
def parse(request):
|
|
return request.param
|
|
|
|
|
|
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
|
def test_no_scheduling(step, parse):
|
|
p = parse(r"This is a (basic:0.6) (prompt) with [no scheduling] features and \(escaped parens\)")
|
|
expected = prompt(1.0, r"This is a (basic:0.6) (prompt) with [no scheduling] features and \(escaped parens\)")
|
|
assert prompts_match(p.at_step(step), expected)
|
|
|
|
|
|
def test_integer_steps(parse):
|
|
p = parse("[a:b:25]", num_steps=50)
|
|
assert prompts_match(p.at_step(0), prompt(0.5, "a"))
|
|
assert prompts_match(p.at_step(25), prompt(0.5, "a"))
|
|
assert prompts_match(p.at_step(0.5), prompt(0.5, "a"))
|
|
assert prompts_match(p.at_step(0.51), prompt(1.0, "b"))
|
|
assert prompts_match(p.at_step(30), prompt(1.0, "b"))
|
|
|
|
|
|
def test_mixed_steps(parse):
|
|
p = parse("[a:b:25] [c:d:0.25]", num_steps=50)
|
|
assert prompts_match(p.at_step(0), prompt(0.25, "a c"))
|
|
assert prompts_match(p.at_step(25), prompt(0.5, "a d"))
|
|
assert prompts_match(p.at_step(0.5), prompt(0.5, "a d"))
|
|
assert prompts_match(p.at_step(0.51), prompt(1.0, "b d"))
|
|
assert prompts_match(p.at_step(30), prompt(1.0, "b d"))
|
|
|
|
|
|
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
|
def test_quote(step, parse):
|
|
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 prompts_match(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, parse):
|
|
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(parse):
|
|
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.5, 0.5, "This is a (basic:0.6) (prompt) with (very [simple]:1.1) features")
|
|
|
|
|
|
@pytest.mark.parametrize("step", [0, 0.5, 1])
|
|
def test_basic_cornercase(parse, step):
|
|
p = parse("This contains[ an ignored segment in:1] the prompt")
|
|
assert_prompt(p, step, 1.0, "This contains the prompt")
|
|
|
|
|
|
def test_basic_ok(parse):
|
|
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.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, parse):
|
|
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 prompts_match(p.at_step(step), expected)
|
|
|
|
|
|
def test_scheduled_lora(parse):
|
|
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))
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"text",
|
|
[
|
|
"This is a sequence of [SEQ:a:0.2::0.5:c:0.8][SEQ: and x:0.8]",
|
|
"This is a sequence of [[a:[c:0.5]:0.2]::0.8][ and x::0.8]",
|
|
],
|
|
)
|
|
def test_seq(parse, text):
|
|
p = parse(text)
|
|
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 ",
|
|
}
|
|
for k, v in prompts.items():
|
|
assert_prompt(p, k, k, v)
|
|
|
|
|
|
def test_shortcuts_scheduling(parse):
|
|
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, parse):
|
|
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, parse):
|
|
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(parse):
|
|
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 prompts_match(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 prompts_match(p2.at_step(0), p2.at_step(1))
|
|
|
|
|
|
def test_def(parse):
|
|
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
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"text, cases",
|
|
[
|
|
(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")]),
|
|
(r"[a:b \(test\):0.2]", [(0, 0.2, r"a"), (0.25, 1, r"b \(test\)")]),
|
|
],
|
|
)
|
|
def test_escapes(text, cases, parse):
|
|
p = parse(text)
|
|
for step, until, val in cases:
|
|
assert_prompt(p, step, until, val)
|
|
|
|
|
|
# I think these were wrong in the old parser too
|
|
@pytest.mark.xfail(reason="Old parser behaviour, possibly buggy")
|
|
@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]")]),
|
|
],
|
|
)
|
|
def test_escapes_fail(text, cases, parse):
|
|
p = parse(text)
|
|
for step, until, val in cases:
|
|
assert_prompt(p, step, until, val)
|
|
|
|
|
|
def test_comments(parse):
|
|
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(parse):
|
|
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
|
|
|
|
|
|
def test_filters(parse):
|
|
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")
|
|
|
|
|
|
def test_emb(parse):
|
|
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_defaultstep(parse):
|
|
p = parse("[cat|dog|tiger]")
|
|
p2 = parse("[cat|dog|tiger:0.1]")
|
|
assert p.parsed_prompt == p2.parsed_prompt
|
|
|
|
|
|
def test_alternating_basic(parse):
|
|
p = parse("[cat|dog|tiger]")
|
|
p2 = parse("[cat|dog|tiger:0.1]")
|
|
assert p.parsed_prompt == p2.parsed_prompt
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"equivalent",
|
|
[
|
|
"[cat::0.1][dog:0.1,0.2][tiger:0.2,0.3][cat:0.3,0.4][dog:0.4,0.5][tiger:0.5,0.6][cat:0.6,0.7][dog:0.7,0.8][tiger:0.8,0.9][cat:0.9,1.0]"
|
|
],
|
|
)
|
|
def test_alternating_equivalences(parse, equivalent):
|
|
p = parse("[cat|dog|tiger]")
|
|
p2 = parse(equivalent)
|
|
assert p.parsed_prompt == p2.parsed_prompt
|
|
|
|
|
|
@pytest.mark.xfail(reason="Old parser behaviour")
|
|
def test_cornercase_failure(parse):
|
|
"""p1 returns a prompt entry until 0 at the start"""
|
|
p = parse("[cat:0,0.1]")
|
|
p2 = parse("[cat::0.1]")
|
|
assert p.parsed_prompt == p2.parsed_prompt
|
|
|
|
|
|
def test_cornercase_corrected(parse):
|
|
p = parse("[cat:0,0.1]")
|
|
p2 = parse("[cat::0.1]")
|
|
assert p.parsed_prompt[0][0] == 0.0
|
|
assert p.parsed_prompt[1:] == p2.parsed_prompt
|
|
|
|
|
|
def test_floats(parse):
|
|
p = parse("[a:b:0.5] [c:d:e:0.2,0.7] <lora:test:-0.3>")
|
|
p2 = parse("[a:b:.5] [c:d:e:.2,.7] <lora:test:-.3>")
|
|
assert p.parsed_prompt == p2.parsed_prompt
|
|
|
|
|
|
def test_alternating_lora(parse):
|
|
p4 = parse("[cat|[dog:wolf<lora:canine:1>:0.5]:0.2]")
|
|
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))
|
|
|
|
|
|
def test_alternating_nested(parse):
|
|
p3 = parse("[cat|[dog|wolf]|tiger]")
|
|
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)
|