diff --git a/prompt_control/prompts.py b/prompt_control/prompts.py index 689ae6c..b51b3c8 100644 --- a/prompt_control/prompts.py +++ b/prompt_control/prompts.py @@ -604,7 +604,7 @@ def encode_prompt(clip, text, start_pct, end_pct, defaults, masks): return f"MASK({args[0]})" for prompt in prompts: - text, noise_w, generator = get_noise(text) + prompt, noise_w, generator = get_noise(prompt) base_prompt, attn_couple_prompts = split_by_function(prompt, "COUPLE", defaults=None, require_args=False) prompts = [base_prompt] + [couple_mask(f.args) + chunk for (chunk, f) in attn_couple_prompts] diff --git a/tests/test_encode.py b/tests/test_encode.py index 08fade6..aa44c2b 100644 --- a/tests/test_encode.py +++ b/tests/test_encode.py @@ -73,6 +73,16 @@ def tensors_equal(t1, t2): npt.assert_equal(t1.detach().numpy(), t2.detach().numpy()) +def cond_neq(c1, c2, key=None, key_assert=None): + ok = False + try: + cond_equal(c1, c2, key=key, key_assert=key_assert) + except AssertionError: + ok = True + if not ok: + raise ValueError("Tensors should not be equal") + + def cond_equal(c1, c2, key=None, key_assert=None): assert len(c1) == len(c2) for i in range(len(c1)): @@ -241,3 +251,15 @@ class TestPCTextEncode: (c2,) = run(pc_text_encode, clip, "test COUPLE MASK(0 0.2, 0.5) prompt1") cond_equal(c, c2) cond_equal(c, c2, "hooks", compare_hookgroup_mask) + + def test_noise_weight0(self, text_encoder_clips, pc_text_encode, node_class_objs): + for _k, clip in text_encoder_clips: + (c1,) = run(pc_text_encode, clip, "test") + (c2,) = run(pc_text_encode, clip, "test NOISE(0, 0)") + cond_equal(c1, c2) + + def test_noise(self, text_encoder_clips, pc_text_encode, node_class_objs): + for _k, clip in text_encoder_clips: + (c1,) = run(pc_text_encode, clip, "test") + (c2,) = run(pc_text_encode, clip, "test NOISE(1, 0)") + cond_neq(c1, c2)