From f4e57ec51429a23b7edaa6c91a370baa2eb01f57 Mon Sep 17 00:00:00 2001 From: asagi4 <130366179+asagi4@users.noreply.github.com> Date: Sun, 8 Jun 2025 00:33:40 +0300 Subject: [PATCH] Switch on new implementation by default --- prompt_control/adv_encode.py | 16 ++++++++++++---- prompt_control/prompts.py | 2 +- 2 files changed, 13 insertions(+), 5 deletions(-) diff --git a/prompt_control/adv_encode.py b/prompt_control/adv_encode.py index c7a2801..5c3c1cf 100644 --- a/prompt_control/adv_encode.py +++ b/prompt_control/adv_encode.py @@ -547,12 +547,20 @@ def advanced_encode_from_tokens( tokenizer=None, **extra_args, ): - if "new+" in weight_interpretation: - weight_interpretation = weight_interpretation.replace("new+", "") + if "old+" not in weight_interpretation: enc = AdvancedEncoder( encode_func, weight_interpretation, token_normalization, tokenizer, m_token, w_max, **extra_args ) - log.info("Using new implementation for %s", weight_interpretation) return enc(tokenized, return_pooled=return_pooled, apply_to_pooled=apply_to_pooled) else: - return old_advanced_encode_from_tokens(tokenized, token_normalization, weight_interpretation, encode_func, 266, return_pooled=return_pooled, apply_to_pooled=apply_to_pooled) + weight_interpretation = weight_interpretation.replace("old+", "") + log.warning("Using old implementation of %s", weight_interpretation) + return old_advanced_encode_from_tokens( + tokenized, + token_normalization, + weight_interpretation, + encode_func, + 266, + return_pooled=return_pooled, + apply_to_pooled=apply_to_pooled, + ) diff --git a/prompt_control/prompts.py b/prompt_control/prompts.py index 8cc843d..018726e 100644 --- a/prompt_control/prompts.py +++ b/prompt_control/prompts.py @@ -65,7 +65,7 @@ def get_style(text, default_style="comfy", default_normalization="none"): style, normalization = styles[0] style = style.strip() normalization = normalization.strip() - if style.replace("new+", "") not in AVAILABLE_STYLES: + if style.replace("old+", "") not in AVAILABLE_STYLES: log.warning("Unrecognized prompt style: %s. Using %s", style, default_style) for part in normalization.split("+"):