diff --git a/docs/SYNTAX.md b/docs/SYNTAX.md index 5f11c67..5867668 100644 --- a/docs/SYNTAX.md +++ b/docs/SYNTAX.md @@ -1,5 +1,9 @@ # Prompt PostProcessor syntax +## Basic usage + +The extension works by modifying the prompt and negative prompt, expanding, replacing and cleaning content, before the image is generated. + ## Commands The extension uses a format for its commands similar to an extranetwork, but it has a `ppp:` prefix followed by the command, and then a space and any parameters (if any). @@ -212,28 +216,33 @@ All these variables can be used to output content or behave differently based on Names starting with an underscore are reserved for system variables: -| System variable | Value | -| --------------- | ----- | -| `_model` | the model identifier (`sd1`, `sd2`, `sdxl`, `sd3`, `flux`, `auraflow`). `_sd` also works but is deprecated. | -| `_modelname` | the model filename (without path). Do not confuse with the `modelname` input in *ComfyUI* which matches actually to the `_modelfullname` variable. `_sdname` also works but is deprecated. | -| `_modelfullname` | the model filename (with path). `_sdfullname` also works but is deprecated. In *ComfyUI* this variable can also be **set** to override the filename used for model detection (see below). | -| `_modelclass` | the class used for the model. Note that this is dependent on the webui. In A1111 all SD versions use the same class. Can be used for new models that are not supported yet with the `_is_*` variables. The debug setting will show all system variables when generating in case you need to see which one to use for a certain model. | -| `_is_kkkk` | true if the model is of kind *kkkk* (the model identifier, f.e. sdxl; those set in the ppp_config.yaml file) | -| `_is_vvvv` | true if the model matches the *vvvv* model variant definition (based on its filename). Note that the corresponding variable for the model kind will also be true. | -| `_is_pure_kkkk` | true if the model is of kind *kkkk* and not a variant. | -| `_is_variant_kkkk` | true if the model version is any variant of model kind *kkkk* and not the pure version. Note that the corresponding variable for the model kind will also be true. | -| `_is_sd` | true if the model is any version of SD | -| `_is_ssd` | true if the model 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 model is SDXL and not an SSD model. | -| `_is_sdxl_no_pony` | true if the model is SDXL and not a Pony model (the `pony` variant must be defined in settings). Kept to maintain compatibility with previous versions. | -| `_opt_...` | All the options. | -| `_input_seed` | The seed used. | -| `_input_pos_prompt` | The original positive prompt. | -| `_input_neg_prompt` | The original negative prompt. | +| System variable | Value | +| --------------- | ----- | +| `_model` | the model identifier (`sd1`, `sd2`, `sdxl`, `sd3`, `flux`, `auraflow`). `_sd` also works but is deprecated. | +| `_modelname` | the model filename (without path). Do not confuse with the `modelname` input in *ComfyUI* which matches actually to the `_modelfullname` variable. `_sdname` also works but is deprecated. | +| `_modelfullname` | the model filename (with path). `_sdfullname` also works but is deprecated. In *ComfyUI* this variable can also be **set** to override the filename used for model detection (see below). | +| `_modelclass` | the class used for the model. Note that this is dependent on the webui. In A1111 all SD versions use the same class. Can be used for new models that are not supported yet with the `_is_*` variables. The debug setting will show all system variables when generating in case you need to see which one to use for a certain model. | +| `_is_kkkk` | true if the model is of kind *kkkk* (the model identifier, f.e. sdxl; those set in the ppp_config.yaml file) | +| `_is_vvvv` | true if the model matches the *vvvv* model variant definition (based on its filename). Note that the corresponding variable for the model kind will also be true. | +| `_is_pure_kkkk` | true if the model is of kind *kkkk* and not a variant. | +| `_is_variant_kkkk` | true if the model version is any variant of model kind *kkkk* and not the pure version. Note that the corresponding variable for the model kind will also be true. | +| `_is_sd` | true if the model is any version of SD | +| `_is_ssd` | true if the model 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 model is SDXL and not an SSD model. | +| `_is_sdxl_no_pony` | true if the model is SDXL and not a Pony model (the `pony` variant must be defined in settings). Kept to maintain compatibility with previous versions. | +| `_opt_...` | All the options. | +| `_input_seed` | The seed used. | +| `_input_pos_prompt` | The original positive prompt. | +| `_input_neg_prompt` | The original negative prompt. | +| `_input_prev_pos_prompt` | The positive prompt result of the previous phase. Available only in the hires fix phase in A1111 compatible hosts and with no combinatorial generation. | +| `_input_prev_neg_prompt` | The negative prompt result of the previous phase. Available only in the hires fix phase in A1111 compatible hosts and with no combinatorial generation. | > [!NOTE] > The model path is relative to the checkpoint/difussion_models folder, just as it appears in the load nodes. +> [!NOTE] +> You can use the `_input_prev_pos_prompt` and `_input_prev_neg_prompt` only in the hires fix prompt boxes in A1111 compatible hosts, and only without the combinatorial option. + ## Set command This command sets the value of a variable that can be checked later. diff --git a/ppp.py b/ppp.py index a967651..a534ae6 100644 --- a/ppp.py +++ b/ppp.py @@ -709,7 +709,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in vs.set_system("_is_pure_" + x, sdchecks[x] and not any(is_models.values())) vs.set_system("_is_variant_" + x, sdchecks[x] and any(is_models.values())) # special cases - vs.set_system("_is_sd", sdchecks.get("sd1", False) or sdchecks.get("sd2", False) or sdchecks.get("sdxl", False) or sdchecks.get("sd3", False)) + vs.set_system( + "_is_sd", + sdchecks.get("sd1", False) + or sdchecks.get("sd2", False) + or sdchecks.get("sdxl", False) + or sdchecks.get("sd3", False), + ) is_ssd = self.state.env_info.get("is_ssd", False) vs.set_system("_is_ssd", is_ssd) vs.set_system("_is_sdxl_no_ssd", sdchecks.get("sdxl", False) and not is_ssd) @@ -1024,6 +1030,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in negative_prompt: str, seed: int, jobinfo: Any = None, + input_vars: dict[str, Any] | None = None, ) -> list[tuple[str, str, dict[str, Any]]]: """ Process the prompt and negative prompt. @@ -1032,6 +1039,8 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in prompt (str): The prompt. negative_prompt (str): The negative prompt. seed (int): The seed for the random number generator. + jobinfo (Any): Additional job information to be stored in the input state. + input_vars (dict[str, Any] | None): Additional input variables to be set as system variables with the "_input_" prefix. Returns: list[tuple[str, str, dict[str, Any]]]: A list of tuples, each containing the processed prompt, negative prompt, and all variables. @@ -1048,6 +1057,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in self.state.inputs.jobinfo = jobinfo # Input related system variables + if input_vars: + for k, v in input_vars.items(): + self.state.variables.set_system("_input_" + k, v) for input_name in self.state.inputs.__dict__.keys(): input_value = getattr(self.state.inputs, input_name) var_name = "_input_" + input_name @@ -1226,6 +1238,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in original_negative_prompt: str, seed: int = -1, jobinfo: Any = None, + input_vars: dict[str, Any] | None = None, ) -> list[tuple[str, str, dict[str, Any]]]: """ Initializes the random number generator and processes the prompt and negative prompt. @@ -1235,6 +1248,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in original_negative_prompt (str): The original negative prompt. seed (int): The seed. jobinfo (Any): Optional job information, available as `_input_jobinfo`. + input_vars (dict[str, Any] | None): Optional dictionary of input variables to set before processing. Returns: list[tuple[str, str, dict[str, Any]]]: A list of tuples containing the processed prompt, negative prompt and all the prompt variables. @@ -1249,7 +1263,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in if self.state.cyclical_state.last_prompt_pair != (original_prompt, original_negative_prompt): self.state.cyclical_state.reset() self.state.cyclical_state.last_prompt_pair = (original_prompt, original_negative_prompt) - results = self.__processprompts(prompt, negative_prompt, seed, jobinfo) + results = self.__processprompts(prompt, negative_prompt, seed, jobinfo, input_vars or {}) t2 = time.monotonic_ns() self.log(logging.INFO, f"Process prompt pair time: {(t2 - t1) / 1_000_000_000:.3f} seconds") # self.log(logging.DEBUG,f"Wildcards memory usage: {self.state.wildcards_obj.__sizeof__()}") diff --git a/scripts/ppp_script.py b/scripts/ppp_script.py index 28073d0..a9e9bd5 100644 --- a/scripts/ppp_script.py +++ b/scripts/ppp_script.py @@ -79,6 +79,7 @@ class PromptPostProcessorA1111Script(scripts.Script): try: # Support for SD.Next import installer # type: ignore + if hasattr(installer, "control_extensions"): if self.title() not in installer.control_extensions: installer.control_extensions.append(self.title()) # We add the extension to the whitelist. @@ -86,7 +87,12 @@ class PromptPostProcessorA1111Script(scripts.Script): # log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.WARNING, "Could not import control_extensions from installer, SD.Next support will not work.") pass except Exception as e: # pylint: disable=broad-except - log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.ERROR, f"Error while adding to the SD.Next extension whitelist: {e}") + log( + self.ppp_logger, + DEBUG_LEVEL.minimal, + logging.ERROR, + f"Error while adding to the SD.Next extension whitelist: {e}", + ) def title(self): """ @@ -388,8 +394,8 @@ class PromptPostProcessorA1111Script(scripts.Script): else: calculated_seeds = seeds - # (prompt type, typeindex) -> (new positive prompt, new negative prompt) - prompts_list: dict[tuple[str, int], tuple[str, str]] = {} + # [index][indextype] -> (new positive prompt, new negative prompt) + prompts_list: list[list[tuple[str, str]]] = [] extra_params = {} # adds prompts @@ -402,10 +408,10 @@ class PromptPostProcessorA1111Script(scripts.Script): rnh: list[str] = getattr(p, "all_hr_negative_prompts", None) hiresfix_exists = bool(rph) and bool(rnh) for i in range(len(calculated_seeds)): - if regular_exists: - prompts_list[(regular_type, i)] = None + prompts_list.append([]) + prompts_list[i].append(None) if hiresfix_exists: - prompts_list[(hiresfix_type, i)] = None + prompts_list[i].append(None) ppp.process_prompts_group_start() if input_combinatorial: @@ -429,7 +435,7 @@ class PromptPostProcessorA1111Script(scripts.Script): num_comb = len(comb_results) for i in range(len(rpr)): # pylint: disable=consider-using-enumerate posp, negp, _ = comb_results[i % num_comb] - prompts_list[(regular_type, i)] = (posp, negp) + prompts_list[i][0] = (posp, negp) extra_params["PPP combination"] = [str(1 + (i % num_comb)) for i in range(len(rpr))] if hiresfix_exists: hiresfix_equal = regular_exists and rph == rpr and rnh == rnr @@ -441,7 +447,7 @@ class PromptPostProcessorA1111Script(scripts.Script): "hiresfix prompts are the same as regular prompts, skipping combinatorial processing for hiresfix", ) for i in range(len(rph)): # pylint: disable=consider-using-enumerate - prompts_list[(hiresfix_type, i)] = prompts_list.get((regular_type, i)) + prompts_list[i][1] = prompts_list[i][0] else: log( self.ppp_logger, @@ -462,43 +468,54 @@ class PromptPostProcessorA1111Script(scripts.Script): num_comb_hr = len(comb_results_hr) for i in range(len(rph)): # pylint: disable=consider-using-enumerate posp, negp, _ = comb_results_hr[i % num_comb_hr] - prompts_list[(hiresfix_type, i)] = (posp, negp) + prompts_list[i][1] = (posp, negp) extra_params["PPP HR combination"] = [str(1 + (i % num_comb_hr)) for i in range(len(rph))] else: # processes prompts - for prompttype, typeindex in prompts_list.keys(): - log( - self.ppp_logger, - self.ppp_debug_level, - logging.INFO, - f"processing prompts ({prompttype}[{typeindex+1}])", - ) - key = ( - (hash_fullenv, calculated_seeds[typeindex], rpr[typeindex], rnr[typeindex]) - if prompttype == regular_type - else (hash_fullenv, calculated_seeds[typeindex], rph[typeindex], rnh[typeindex]) - ) - cached = self.lru_cache.get(key) - if cached is None: - hsh, seed, prompt, negative_prompt = key - results = ppp.process_prompt( - prompt, - negative_prompt, - seed, - jobinfo={ - "job_timestamp": shared.state.job_timestamp, - "job": shared.state.job, - "detail": f"{prompttype} prompt", - }, + for index, grouplist in enumerate(prompts_list): + for typeindex in range(len(grouplist)): + typeprompt = [regular_type, hiresfix_type][typeindex] + log( + self.ppp_logger, + self.ppp_debug_level, + logging.INFO, + f"processing prompts ({typeprompt}[{index+1}])", ) - posp, negp, _ = results[0] - cached = (posp, negp) - self.lru_cache.put(key, cached) - # adds also the result so i2i doesn't process it unnecessarily - self.lru_cache.put((hsh, seed, posp, negp), cached) - else: - log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache") - prompts_list[(prompttype, typeindex)] = cached + key = ( + (hash_fullenv, calculated_seeds[index], rpr[index], rnr[index]) + if typeindex == 0 + else (hash_fullenv, calculated_seeds[index], rph[index], rnh[index]) + ) + cached = self.lru_cache.get(key) + if cached is None: + hsh, seed, prompt, negative_prompt = key + if typeindex > 0: + prev_prompts = prompts_list[index][typeindex - 1] + input_vars = { + "prev_pos_prompt": prev_prompts[0], + "prev_neg_prompt": prev_prompts[1], + } + else: + input_vars = None + results = ppp.process_prompt( + prompt, + negative_prompt, + seed, + jobinfo={ + "job_timestamp": shared.state.job_timestamp, + "job": shared.state.job, + "detail": f"{typeprompt} prompt", + }, + input_vars=input_vars, + ) + posp, negp, _ = results[0] + cached = (posp, negp) + self.lru_cache.put(key, cached) + # adds also the result so i2i doesn't process it unnecessarily + self.lru_cache.put((hsh, seed, posp, negp), cached) + else: + log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache") + prompts_list[index][typeindex] = cached ppp.process_prompts_group_end() # updates the prompts @@ -506,17 +523,21 @@ class PromptPostProcessorA1111Script(scripts.Script): hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None) regular_changes = False hiresfix_changes = False - for (prompttype, typeindex), (posp, negp) in prompts_list.items(): - if prompttype == regular_type: - if rpr[typeindex].strip() != posp.strip() or rnr[typeindex].strip() != negp.strip(): - regular_changes = True - rpr[typeindex] = posp - rnr[typeindex] = negp - elif prompttype == hiresfix_type: - if rph[typeindex].strip() != posp.strip() or rnh[typeindex].strip() != negp.strip(): - hiresfix_changes = True - rph[typeindex] = posp - rnh[typeindex] = negp + for index, grouplist in enumerate(prompts_list): + for typeindex, groupprompts in enumerate(grouplist): + if groupprompts is None: + continue + (posp, negp) = groupprompts + if typeindex == 0: + if rpr[index].strip() != posp.strip() or rnr[index].strip() != negp.strip(): + regular_changes = True + rpr[index] = posp + rnr[index] = negp + elif typeindex == 1: + if rph[index].strip() != posp.strip() or rnh[index].strip() != negp.strip(): + hiresfix_changes = True + rph[index] = posp + rnh[index] = negp # initialize extra generation parameters if add_prompts: diff --git a/tests/base_tests.py b/tests/base_tests.py index 4b78bfd..f18e0b5 100644 --- a/tests/base_tests.py +++ b/tests/base_tests.py @@ -204,6 +204,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): combinatorial_limit: int = 0, specific_wc_folders: Optional[list[Path]] = None, specific_em_folders: Optional[list[Path]] = None, + input_vars: Optional[dict[str, Any]] = None, ): """ Process the prompt and compare the results with the expected prompts. @@ -218,6 +219,7 @@ class TestPromptPostProcessorBase(unittest.TestCase): combinatorial_limit (int, optional): The combinatorial limit. Defaults to 0. specific_wc_folders (Optional[list[Path]], optional): A list of specific wildcard folders to refresh. Defaults to None. specific_em_folders (Optional[list[Path]], optional): A list of specific extranetwork mapping folders to refresh. Defaults to None. + input_vars (Optional[dict[str, Any]], optional): A dictionary of input variables. Defaults to None. Returns: None @@ -246,6 +248,8 @@ class TestPromptPostProcessorBase(unittest.TestCase): input_prompts.prompt, input_prompts.negative_prompt, seed, + jobinfo={"test_case": self.id()}, + input_vars=input_vars, ) the_obj.process_prompts_group_end() self.assertTrue( @@ -311,6 +315,8 @@ class TestPromptPostProcessorBase(unittest.TestCase): input_prompts.prompt, input_prompts.negative_prompt, seed, + jobinfo={"test_case": self.id()}, + input_vars=input_vars, ) self.assertTrue( self.interrupted == interrupted, diff --git a/tests/tests_varcomms.py b/tests/tests_varcomms.py index 62cb1ed..ec7a765 100644 --- a/tests/tests_varcomms.py +++ b/tests/tests_varcomms.py @@ -766,6 +766,21 @@ class TestVarCommands(TestPromptPostProcessorBase): ), ) + # Input variables + + def test_input_variables(self): + self.process( + InputTuple( + "${_input_prev_positive_prompt}, high quality", + "", + ), + OutputTuple( + "this is a test, high quality", + "", + ), + input_vars={"prev_positive_prompt": "this is a test"}, + ) + # Command tests def test_cmd_stn_complex_features(self): # complex stn command with AND, BREAK and other features