import pytest 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 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)) params = [] params.append(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 ") 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 [::0.3] " ) 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::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:0.6]:0.5]:HR]") p2 = parse("test [:[a:[:b:0.6]:0.5]:HR]") assert p.parsed_prompt == p2.parsed_prompt def test_filters(parse): p = parse("test [[a:[b:0.6]:0.5]:HR]") p2 = parse("test [:[a:[:b: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("[:[: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 [::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_ltgt_in_schedule(parse): p = parse("This should [ correctly:be :0.1]") assert_prompt(p, 0.1, 0.1, "This should correctly", ("test", 1.0, 1.0)) assert_prompt(p, 0.15, 1.0, "This should be ", ("test", 1.0, 1.0)) def test_floats(parse): p = parse("[a:b:0.5] [c:d:e:0.2,0.7] ") p2 = parse("[a:b:.5] [c:d:e:.2,.7] ") assert p.parsed_prompt == p2.parsed_prompt def test_alternating_lora(parse): p4 = parse("[cat|[dog:wolf: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)