From b06bcfa677c5c86851e70406b57cd1ad626cee16 Mon Sep 17 00:00:00 2001 From: Antonio Cordero Balcazar Date: Sun, 15 Dec 2024 13:32:05 +0100 Subject: [PATCH] * Turned pony detection into detection of any kind of variant models based on filename. Added illustrious. Breaking change: the setting is different and changed pony substrings are not imported. * Fixed and refactored some UI settings. Slight change in the generation of batched seeds. * Detection of reforge. * Small refactoring when getting choices. --- README.md | 44 ++++++++---- ppp.py | 121 ++++++++++++++++++++++----------- ppp_comfyui.py | 17 ++--- pyproject.toml | 6 +- scripts/ppp_script.py | 153 +++++++++++++++++++++++++++++++----------- tests/tests.py | 26 ++++++- 6 files changed, 260 insertions(+), 107 deletions(-) diff --git a/README.md b/README.md index ea183ae..9ab2e89 100644 --- a/README.md +++ b/README.md @@ -54,9 +54,8 @@ On SD.Next I recommend you disable the native wildcard processing. On ComfyUI: 1. Go to Manager > Custom Nodes Manager -2. Install through ComfyUI Manager -3. Click Install via Git URL and enter -4. Restart +2. Search for "Prompt PostProcessor" and install or click Install via Git URL and enter +3. Restart ## Usage @@ -118,7 +117,7 @@ These are examples of formats you can use to insert a choice construct: Notes: * The Dynamic Prompts format `{2$$__flavours__}` does not work as expected. It will only output one value. You can write is as `{r2$$__flavours__}` to get two values, but they may repeat since the evaluation of the wildcard is independent of the choices selection. -* Whitespace in the choices is not ignored like in Dynamic Prompts, but will be cleaned up if the appropriate settings are checked. +* Whitespace around the choices is not ignored like in Dynamic Prompts, but will be cleaned up if the appropriate settings are checked. ### Wildcards @@ -246,6 +245,8 @@ The full format is: content onecontent twoother content ``` +Any `elif`s (there can be multiple) and the `else` are optional. + The *conditionN* compares a variable with a value or a list of values. The allowed formats are: ```text @@ -272,15 +273,24 @@ The variable can be one set with the `set` or `add` commands or you can use inte * `_is_sd1`: true if the loaded model version is SD 1.x * `_is_sd2`: true if the loaded model version is SD 2.x * `_is_sdxl`: true if the loaded model version is SDXL (includes Pony models) -* `_is_ssd`: true if the loaded model version is SSD (Segmind Stable Diffusion 1B). Note that for an SSD model `_is_sdxl` will also be true. -* `_is_sdxl_no_ssd`: true if the loaded model version is SDXL and not an SSD model. -* `_is_pony`: true if the loaded model version is SDXL and a Pony model (based on its filename). Note that for a pony model `_is_sdxl` will also be true. -* `_is_sdxl_no_pony`: true if the loaded model version is SDXL and not a Pony model. * `_is_sd3`: true if the loaded model version is SD 3.x * `_is_flux`: true if the loaded model is Flux * `_is_auraflow`: true if the loaded model is AuraFlow +* `_is_ssd`: true if the loaded model version is SSD (Segmind Stable Diffusion 1B). Note that for an SSD model `_is_sdxl` will also be true. +* `_is_sdxl_no_ssd`: true if the loaded model version is SDXL and not an SSD model. -Any `elif`s (there can be multiple) and the `else` are optional. +Then there are also variables for the user defined model variants defined by the "model variant definitions" setting. This is where the pony, and now also illustrious, definitions are to detect those models. + +* `_is_xxxx`: true if the loaded model matches the xxxx definition (based on its filename). Note that the corresponding variable for the model kind will also be true. + +To maintain compatibility with previous versions the following variable still exists: + +* `_is_sdxl_no_pony`: true if the loaded model version is SDXL and not a Pony model (the "pony" variant must be defined in settings). + +But in general these new variables are created for all model types: + +* `_is_pure_xxxx`: true if the loaded model is of kind xxxx (f.e. sdxl) and not a variant. +* `_is_variant_xxxx`: true if the loaded model version is any variant of model kind xxxx and not the pure version. #### Example @@ -289,7 +299,7 @@ Any `elif`s (there can be multiple) and the `else` are optional. ```text test sd1x test pony - test sdxl + test sdxl unknown model ``` @@ -362,9 +372,9 @@ This should still work as intended, and the only negative point i see is the unn ### A1111 (and compatible UIs) UI options * **Force equal seeds**: Changes the image seeds and variation seeds to be equal to the first of the batch. This allows using the same values for all the images in a batch. -* **Unlink seed**: Uses the specified seed for the prompt generation instead of the one from the image. -* **Seed**: The seed to use for the prompt generation. If -1 a random one will be used for each image in the batch. This seed is only used for wildcards and choices. -* **Variable seed**: If the seed is not -1 you can use this to increase it for the other images in the batch. +* **Unlink seed**: Uses the specified seed for the prompt generation instead of the one from the image. This seed is only used for wildcards and choices. +* **Prompt seed**: The seed to use for the prompt generation. If -1 a random one will be used. +* **Incremental seed**: When using a batch you can use this to set the rest of the prompt seeds with consecutive values. ### ComfyUI specific inputs @@ -377,7 +387,13 @@ This should still work as intended, and the only negative point i see is the unn ### General settings * **Debug level**: what to write to the console. Note: in SD.Next debug messages only show if you launch it with the --debug argument. -* **Pony substrings**: list of substrings to detect a Pony model. +* **Model variant definitions**: definitions for model variants to be recognized based on strings found in the full filename. + + The format for each line is (with *kind* being one of the base model identifiers or not defined): + + ```name(kind)=comma separated list of substrings (case insensitive)``` + + The default value defines strings for Pony and Illustrious models. * **Apply in img2img**: check if you want to do the processing in img2img processes (does not apply to ComfyUI node). ### Wildcard settings diff --git a/ppp.py b/ppp.py index 3edf6c8..6ca523b 100644 --- a/ppp.py +++ b/ppp.py @@ -1,4 +1,3 @@ -from functools import reduce import logging import math import os @@ -52,12 +51,21 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in stop = "stop" DEFAULT_STN_SEPARATOR = ", " - DEFAULT_PONY_SUBSTRINGS = ",".join(["pony", "pny", "pdxl"]) + DEFAULT_VARIANTS_DEFINITIONS = "pony(sdxl)=pony,pny,pdxl\nillustrious(sdxl)=illustrious,ilxl" DEFAULT_CHOICE_SEPARATOR = ", " WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK ' WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK " UNPROCESSED_STOP = "UNPROCESSED CONSTRUCTS!\nBREAK " + SUPPORTED_MODELS = [ + "sd1", + "sd2", + "sdxl", + "sd3", + "flux", + "auraflow", + ] + def __init__( self, logger: logging.Logger, @@ -87,9 +95,27 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in # General options self.debug_level = DEBUG_LEVEL(options.get("debug_level", DEBUG_LEVEL.none.value)) - self.pony_substrings = list( - x.strip() for x in (str(options.get("pony_substrings", self.DEFAULT_PONY_SUBSTRINGS))).split(",") - ) + variants_definitions_option = str(options.get("variants_definitions", self.DEFAULT_VARIANTS_DEFINITIONS)) + self.variants_definitions = {} + if variants_definitions_option: + lines = variants_definitions_option.splitlines() + for line in lines: + if "=" in line: + model_tag, elements = line.split("=", 1) + model_name, model_type = re.match(r"(\w+)(?:\((\w+)\))?", model_tag).groups() + if model_type is not None and model_type not in self.SUPPORTED_MODELS: + self.logger.warning( + f"Unsupported model type '{model_type}' in definition for variant '{model_name}'." + ) + elif model_name in self.SUPPORTED_MODELS: + self.logger.warning( + f"Invalid model name in definition for variant '{model_name}'." + ) + else: + self.variants_definitions[model_name.strip()] = ( + model_type or "", + [element.strip() for element in elements.split(",")], + ) # Wildcards options self.wil_process_wildcards = options.get("process_wildcards", True) self.wil_keep_choices_order = options.get("keep_choices_order", False) @@ -179,36 +205,37 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in Initializes the system variables. """ self.system_variables = {} - sdchecks = { - "sd1": self.env_info.get("is_sd1", False), - "sd2": self.env_info.get("is_sd2", False), - "sdxl": self.env_info.get("is_sdxl", False), - "sd3": self.env_info.get("is_sd3", False), - "flux": self.env_info.get("is_flux", False), - "auraflow": self.env_info.get("is_auraflow", False), - "": True, - } + sdchecks = {x: self.env_info.get("is_" + x, False) for x in self.SUPPORTED_MODELS} + sdchecks.update({"": True}) self.system_variables["_model"] = [k for k, v in sdchecks.items() if v][0] self.system_variables["_sd"] = self.system_variables["_model"] # deprecated model_filename = self.env_info.get("model_filename", "") - is_pony = any(s in model_filename.lower() for s in self.pony_substrings) - is_ssd = self.env_info.get("is_ssd", False) self.system_variables["_sdfullname"] = model_filename # deprecated self.system_variables["_modelfullname"] = model_filename self.system_variables["_sdname"] = os.path.basename(model_filename) # deprecated self.system_variables["_modelname"] = os.path.basename(model_filename) self.system_variables["_modelclass"] = self.env_info.get("model_class", "") - self.system_variables["_is_sd1"] = sdchecks["sd1"] - self.system_variables["_is_sd2"] = sdchecks["sd2"] - self.system_variables["_is_sdxl"] = sdchecks["sdxl"] + is_models = { + model_name: (model_type_and_substrings[0] == "" or sdchecks.get(model_type_and_substrings[0], False)) + and any(s in model_filename.lower() for s in model_type_and_substrings[1]) + for model_name, model_type_and_substrings in self.variants_definitions.items() + if model_name not in self.SUPPORTED_MODELS + } + self.system_variables.update({"_is_" + x: y for x, y in is_models.items()}) + for x in sdchecks.keys(): + if x != "": + self.system_variables["_is_" + x] = sdchecks[x] + self.system_variables["_is_pure_" + x] = sdchecks[x] and not any(is_models.values()) + self.system_variables["_is_variant_" + x] = sdchecks[x] and any(is_models.values()) + # special cases + self.system_variables["_is_sd"] = sdchecks["sd1"] or sdchecks["sd2"] or sdchecks["sdxl"] or sdchecks["sd3"] + is_ssd = self.env_info.get("is_ssd", False) self.system_variables["_is_ssd"] = is_ssd self.system_variables["_is_sdxl_no_ssd"] = sdchecks["sdxl"] and not is_ssd - self.system_variables["_is_pony"] = sdchecks["sdxl"] and is_pony - self.system_variables["_is_sdxl_no_pony"] = sdchecks["sdxl"] and not is_pony - self.system_variables["_is_sd3"] = sdchecks["sd3"] - self.system_variables["_is_sd"] = sdchecks["sd1"] or sdchecks["sd2"] or sdchecks["sdxl"] or sdchecks["sd3"] - self.system_variables["_is_flux"] = sdchecks["flux"] - self.system_variables["_is_auraflow"] = sdchecks["auraflow"] + # backcompatibility (but the modern one to use would be _is_pure_sdxl) + self.system_variables["_is_sdxl_no_pony"] = sdchecks["sdxl"] and not self.system_variables.get( + "_is_pony", False + ) def __add_to_insertion_points( self, negative_prompt: str, add_at_insertion_point: list[str], insertion_at: list[tuple[int, int]] @@ -729,10 +756,16 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in Returns: bool: The result of the condition evaluation. """ - var_value = self.__ppp.system_variables.get(cond_var, self.__get_user_variable_value(cond_var)) - if var_value is None: - var_value = "" - self.__ppp.logger.warning(f"Unknown variable {cond_var}") + if cond_var.startswith("_"): # system variable + var_value = self.__ppp.system_variables.get(cond_var, None) + if var_value is None: + var_value = "" + self.__ppp.logger.warning(f"Unknown system variable {cond_var}") + else: # user variable + var_value = self.__get_user_variable_value(cond_var) + if var_value is None: + var_value = "" + self.__ppp.logger.warning(f"Unknown user variable {cond_var}") if isinstance(var_value, str): var_value = var_value.lower() if isinstance(cond_value, list): @@ -758,7 +791,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in ( c[1:-1].lower() if c.startswith('"') or c.startswith("'") - else True if c.lower() == "true" else False if c.lower() == "false" else int(c) + else True if c.lower() == "true" else False if c.lower() == "false" or c == "" else int(c) ) for c in cond_value ) @@ -769,7 +802,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if isinstance(c, str) else ( True - if isinstance(c, bool) and var_value != "false" and var_value is not False + if isinstance(c, bool) and var_value != "false" and var_value != "" and var_value is not False else ( False if isinstance(c, bool) and (var_value != "true" or var_value is False) @@ -1154,13 +1187,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in t2 = time.time() self.__debug_end("extranetworktag", start_result, t2 - t1) - def __get_choices( + def __get_choices_internal( self, options: dict | None, choice_values: list[dict], filter_specifier: Optional[list[list[str]]] = None, wildcard_key: str = None, - ) -> str: + ) -> tuple[str, list[str], str, str]: """ Select choices based on the options. @@ -1171,7 +1204,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in wildcard_key (str): The wildcard key if it is a wildcard. Returns: - str: The selected choice. + tuple: A tuple containing the prefix, selected choices, separator and suffix """ if options is None: options = {} @@ -1188,7 +1221,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in msg = f"wildcard '{wildcard_key}'" if wildcard_key else "choices" self.__ppp.logger.warning(f"Unsupported sampler '{sampler}' in {msg} options!") self.__ppp.interrupt() - return "" + return ("", [], separator, "") if filter_specifier is not None: filtered_choice_values = [] for i, c in enumerate(choice_values): @@ -1297,8 +1330,18 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in suffix = " " + suffix # remove comments results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text] - return prefix + separator.join(results) + suffix - return "" + return (prefix, results, separator, suffix) + return ("", "", separator, "") + + def __get_choices( + self, + options: dict | None, + choice_values: list[dict], + filter_specifier: Optional[list[list[str]]] = None, + wildcard_key: str = None, + ) -> str: + r = self.__get_choices_internal(options, choice_values, filter_specifier, wildcard_key) + return r[0] + r[2].join(r[1]) + r[3] def __convert_choices_options(self, options: Optional[lark.Tree]) -> dict: """ @@ -1490,7 +1533,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in t2 = time.time() self.__debug_end("wildcard", start_result, t2 - t1, wc) return - filter_specifier = None + filter_specifier:list[int|str] = None filter_object = tree.children[2] if filter_object is not None: if ( @@ -1515,7 +1558,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if ( len(selected_wildcards) > 1 and filter_specifier is not None - and any(x.isdecimal() for x in reduce(lambda x, y: x + y, filter_specifier)) + and any(x.isdecimal() for x in filter_specifier) ): self.__ppp.logger.warning( f"Using a globbing wildcard '{wildcard_key}' with positional index filters is not recommended!" diff --git a/ppp_comfyui.py b/ppp_comfyui.py index 6a1a8fd..86767fe 100644 --- a/ppp_comfyui.py +++ b/ppp_comfyui.py @@ -95,12 +95,13 @@ class PromptPostProcessorComfyUINode: "forceInput": False, }, ), - "pony_substrings": ( + "variants_definitions": ( "STRING", { - "default": PromptPostProcessor.DEFAULT_PONY_SUBSTRINGS, - "placeholder": "comma separated list", - "tooltip": "Comma separated list of substrings to look for in the modelname to determine if the model is a pony model", + "default": PromptPostProcessor.DEFAULT_VARIANTS_DEFINITIONS, + "multiline": True, + "placeholder": "", + "tooltip": "Definitions for variant models to be recognized based on strings found in the full filename. Format for each line is: 'name(kind)=comma separated list of substrings (case insensitive)' with kind being one of the base model types or not specified", "defaultInput": False, "forceInput": False, }, @@ -333,7 +334,7 @@ class PromptPostProcessorComfyUINode: neg_prompt, seed, debug_level, # pylint: disable=unused-argument - pony_substrings, + variants_definitions, wc_process_wildcards, wc_wildcards_folders, wc_if_wildcards, @@ -361,7 +362,7 @@ class PromptPostProcessorComfyUINode: "pos_prompt": pos_prompt, "neg_prompt": neg_prompt, "seed": seed, - "pony_substrings": pony_substrings, + "variants_definitions": variants_definitions, "process_wildcards": wc_process_wildcards, "wildcards_folders": wc_wildcards_folders, "if_wildcards": wc_if_wildcards, @@ -392,7 +393,7 @@ class PromptPostProcessorComfyUINode: neg_prompt, seed, debug_level, - pony_substrings, + variants_definitions, wc_process_wildcards, wc_wildcards_folders, wc_if_wildcards, @@ -448,7 +449,7 @@ class PromptPostProcessorComfyUINode: ] options = { "debug_level": debug_level, - "pony_substrings": pony_substrings, + "variants_definitions": variants_definitions, "process_wildcards": wc_process_wildcards, "if_wildcards": wc_if_wildcards, "choice_separator": wc_choice_separator, diff --git a/pyproject.toml b/pyproject.toml index bc429d8..0e47adb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,9 +1,9 @@ [project] name = "sd-webui-prompt-postprocessor" description = "Stable Diffusion WebUI & ComfyUI extension to post-process the prompt, including sending content from the prompt to the negative prompt and wildcards." -version = "2.8.1" -license = {file = "LICENSE.txt"} -dependencies = ["lark"] +version = "2.9.0" +license = { file = "LICENSE.txt" } +dependencies = ["lark", "numpy", "pyyaml"] [project.urls] Repository = "https://github.com/acorderob/sd-webui-prompt-postprocessor" diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 2d52c4a..3525ac1 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -37,6 +37,16 @@ class PromptPostProcessorA1111Script(scripts.Script): __on_ui_settings(): Callback function for UI settings. """ + instance_count = 0 + + @classmethod + def increment_instance_count(cls): + cls.instance_count += 1 + + @classmethod + def get_instance_count(cls): + return cls.instance_count + def __init__(self): """ Initializes the PromptPostProcessor object. @@ -49,6 +59,7 @@ class PromptPostProcessorA1111Script(scripts.Script): Returns: None """ + self.increment_instance_count() lf = PromptPostProcessorLogFactory() self.name = PromptPostProcessor.NAME self.ppp_logger = lf.log @@ -58,7 +69,9 @@ class PromptPostProcessorA1111Script(scripts.Script): with open(grammar_filename, "r", encoding="utf-8") as file: self.grammar_content = file.read() self.wildcards_obj = PPPWildcards(lf.log) - self.ppp_logger.info(f"{PromptPostProcessor.NAME} {PromptPostProcessor.VERSION} initialized") + i = self.get_instance_count() + if i == 1: # some UIs create multiple instances + self.ppp_logger.info(f"{PromptPostProcessor.NAME} {PromptPostProcessor.VERSION} initialized") def title(self): """ @@ -86,48 +99,67 @@ class PromptPostProcessorA1111Script(scripts.Script): force_equal_seeds = gr.Checkbox( label="Force equal seeds", info="Force all image seeds and variation seeds to be equal to the first one, disabling the default autoincrease.", - default=False, + value=False, # show_label=True, elem_id="ppp_force_equal_seeds", ) - gr.HTML( - """
Unlink the seed to use the specified one for the prompts instead of the image seed. - This seed will only change for each image in the batch if the value is -1 or 'variable seed' is checked.
-
Seeds are only used for the wildcards and choice constructs.
""" + gr.HTML("
") + gr.Markdown( + """ + Unlink the seed to use the specified one for the prompts instead of the image seed. + + * A seed of -1 and "Incremental seed" checked will use a random seed for the first prompt and consecutive values for the rest. This is the same as when you use -1 for the image seed. + * A seed of -1 and "Incremental seed" unchecked will use a random seed for each prompt. + * Any other seed value and "Incremental seed" checked will use the specified seed for the first prompt and consecutive values for the rest. + * Any other seed value and "Incremental seed" unchecked will use the specified seed for all the prompts. + + Seeds are only used for the wildcards and choice constructs. + """ ) - unlink_seed = gr.Checkbox( - label="Unlink seed", - default=False, - # show_label=True, - elem_id="ppp_unlink_seed", - ) - seed = gr.Number( - label="Seed", - default=-1, - precision=0, - # minimum=-1, - # maximum=2**32 - 1, - # step=1, - # show_label=True, - min_width=100, - elem_id="ppp_seed", - ) - variable_seed = gr.Checkbox( - label="Variable seed", - default=False, - # show_label=True, - elem_id="ppp_variable_seed", - ) - return [force_equal_seeds, unlink_seed, seed, variable_seed] + gr.HTML("
") + with gr.Row(equal_height=True): + unlink_seed = gr.Checkbox( + label="Unlink seed", + value=False, + # show_label=True, + elem_id="ppp_unlink_seed", + ) + seed = gr.Number( + label="Prompt seed", + value=-1, + precision=0, + # minimum=-1, + # maximum=2**32 - 1, + # step=1, + # show_label=True, + min_width=100, + elem_id="ppp_seed", + ) + incremental_seed = gr.Checkbox( + label="Incremental seed (only applies to batches)", + value=False, + # show_label=True, + elem_id="ppp_incremental_seed", + ) + return [force_equal_seeds, unlink_seed, seed, incremental_seed] def process( - self, p: StableDiffusionProcessing, input_force_equal_seeds, input_unlink_seed, input_seed, input_variable_seed + self, + p: StableDiffusionProcessing, + input_force_equal_seeds, + input_unlink_seed, + input_seed, + input_incremental_seed, ): # pylint: disable=arguments-differ """ Processes the prompts and applies post-processing operations. Args: p (StableDiffusionProcessing): The StableDiffusionProcessing object containing the prompts. + input_force_equal_seeds (bool): Flag indicating whether to force equal seeds. + input_unlink_seed (bool): Flag indicating whether to unlink the seed. + input_seed (int): The seed value. + input_incremental_seed (bool): Flag indicating whether to use incremental seed. Returns: None @@ -143,13 +175,35 @@ class PromptPostProcessorA1111Script(scripts.Script): if self.ppp_debug_level != DEBUG_LEVEL.none: self.ppp_logger.info("Not processing the prompt for i2i") return - if self.ppp_debug_level != DEBUG_LEVEL.none: - self.ppp_logger.info(f"Post-processing prompts ({'i2i' if is_i2i else 't2i'})") + app_names = { + "sdnext": "SD.Next", + "forge": "Forge", + "reforge": "reForge", + "a1111": "A1111 (or compatible)", + } app = ( "forge" if hasattr(p.sd_model, "model_config") - else "sdnext" if hasattr(p.sd_model, "is_sdxl") and not hasattr(p.sd_model, "is_ssd") else "a1111" + else ( + "reforge" + if hasattr(p.sd_model, "forge_objects") + else ("sdnext" if hasattr(p.sd_model, "is_sdxl") and not hasattr(p.sd_model, "is_ssd") else "a1111") + ) ) + if self.ppp_debug_level != DEBUG_LEVEL.none: + self.ppp_logger.info(f"Post-processing prompts ({'i2i' if is_i2i else 't2i'}) running on {app_names[app]}") + models_supported = {x: True for x in PromptPostProcessor.SUPPORTED_MODELS} + if app == "sdnext": + models_supported["ssd"] = False + elif app == "forge": + models_supported["ssd"] = False + models_supported["auraflow"] = False + elif app == "reforge": + models_supported["flux"] = False + models_supported["auraflow"] = False + else: # assume A1111 compatible + models_supported["flux"] = False + models_supported["auraflow"] = False env_info = { "app": app, "models_path": models_path, @@ -183,6 +237,15 @@ class PromptPostProcessorA1111Script(scripts.Script): env_info["is_sd3"] = getattr(p.sd_model, "is_sd3", False) env_info["is_flux"] = p.sd_model.model_config.__class__.__name__ == "Flux" env_info["is_auraflow"] = False # p.sd_model.model_config.__class__.__name__ == "AuraFlow" + elif app == "reforge": + env_info["model_class"] = p.sd_model.__class__.__name__ + env_info["is_sd1"] = getattr(p.sd_model, "is_sd1", False) + env_info["is_sd2"] = getattr(p.sd_model, "is_sd2", False) + env_info["is_sdxl"] = getattr(p.sd_model, "is_sdxl", False) + env_info["is_ssd"] = getattr(p.sd_model, "is_ssd", False) + env_info["is_sd3"] = getattr(p.sd_model, "is_sd3", False) + env_info["is_flux"] = False + env_info["is_auraflow"] = False else: # assume A1111 compatible (p.sd_model.__class__.__name__=="DiffusionEngine") env_info["model_class"] = p.sd_model.__class__.__name__ env_info["is_sd1"] = getattr(p.sd_model, "is_sd1", False) @@ -202,7 +265,9 @@ class PromptPostProcessorA1111Script(scripts.Script): ] options = { "debug_level": getattr(opts, "ppp_gen_debug_level", DEBUG_LEVEL.none.value), - "pony_substrings": getattr(opts, "ppp_gen_ponysubstrings", PromptPostProcessor.DEFAULT_PONY_SUBSTRINGS), + "variants_definitions": getattr( + opts, "ppp_gen_variantsdefinitions", PromptPostProcessor.DEFAULT_VARIANTS_DEFINITIONS + ), "process_wildcards": getattr(opts, "ppp_wil_processwildcards", True), "if_wildcards": getattr(opts, "ppp_wil_ifwildcards", PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value), "choice_separator": getattr(opts, "ppp_wil_choice_separator", PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR), @@ -241,10 +306,11 @@ class PromptPostProcessorA1111Script(scripts.Script): if self.ppp_debug_level != DEBUG_LEVEL.none: self.ppp_logger.info("Using unlinked seed") num_seeds = len(getattr(p, "all_seeds", [])) - if input_seed == -1: + if input_incremental_seed: + first_seed = np.random.randint(0, 2**32, dtype=np.int64) if input_seed == -1 else input_seed + calculated_seeds = [first_seed + i for i in range(num_seeds)] + elif input_seed == -1: calculated_seeds = np.random.randint(0, 2**32, size=num_seeds, dtype=np.int64) - elif input_variable_seed: - calculated_seeds = [input_seed + i for i in range(num_seeds)] else: calculated_seeds = [input_seed for _ in range(num_seeds)] else: @@ -379,10 +445,15 @@ def on_ui_settings(): ), ) shared.opts.add_option( - key="ppp_gen_ponysubstrings", + key="ppp_gen_variantsdefinitions", info=shared.OptionInfo( - PromptPostProcessor.DEFAULT_PONY_SUBSTRINGS, - label="Comma separated list of substrings to look for in the model full filename to flag it as Pony (case insensitive)", + PromptPostProcessor.DEFAULT_VARIANTS_DEFINITIONS, + label="Definitions for variant models", + comment_after="Recognized based on strings found in the full filename. Format for each line is: 'name(kind)=comma separated list of substrings (case insensitive)' with kind being one of the base model types (" + + ",".join(PromptPostProcessor.SUPPORTED_MODELS) + + ") or not specified.", + component=gr.Textbox, + component_args={"lines": 7}, section=section, ), ) diff --git a/tests/tests.py b/tests/tests.py index 8245523..34f3e6b 100644 --- a/tests/tests.py +++ b/tests/tests.py @@ -29,7 +29,7 @@ class TestPromptPostProcessor(unittest.TestCase): self.__ppp_logger.setLevel(logging.DEBUG) self.__defopts = { "debug_level": DEBUG_LEVEL.full.value, - "pony_substrings": PromptPostProcessor.DEFAULT_PONY_SUBSTRINGS, + "variants_definitions": PromptPostProcessor.DEFAULT_VARIANTS_DEFINITIONS, "process_wildcards": True, "if_wildcards": PromptPostProcessor.IFWILDCARDS_CHOICES.ignore.value, "choice_separator": ", ", @@ -355,7 +355,7 @@ class TestPromptPostProcessor(unittest.TestCase): def test_cmd_if_nested(self): # nested if command self.__process( PromptPair( - "this is SD1PONYSD2", "" + "this is SD1PONYSD2NOPONYNOPONY", "" ), PromptPair("this is PONY", ""), ppp=PromptPostProcessor( @@ -824,6 +824,28 @@ class TestPromptPostProcessor(unittest.TestCase): ppp=self.__nocupppp, ) + # Model variants tests + + def test_variants(self): + self.__process( + PromptPair("test1test2test3test4", ""), + PromptPair("test1test2", ""), + ppp=PromptPostProcessor( + self.__ppp_logger, + self.__interrupt, + { + **self.__def_env_info, + "model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", + }, + { + **self.__defopts, + "variants_definitions": "test1(sdxl)=testmodel\ntest2=testmodel\ntest3(sd1)=testmodel\ntest4(invalid)=testmodel\nsdxl()=testmodel", + }, + self.__grammar_content, + self.__wildcards_obj, + ), + ) + # ComfyUI tests def test_comfyui_attention(self): # attention conversion