Compare commits

...
19 Commits
Author SHA1 Message Date
Antonio Cordero Balcazar cd3a18e651 * Added MiniMax H3 model 2026-08-14 18:53:47 +02:00
Antonio Cordero Balcazar f8c15d23c3 * Added next_seed option.
* Added output variable `_output_seed`.
* Fixed some table formatting in the documentation.
2026-08-14 18:14:48 +02:00
Antonio Cordero Balcazar 671d8b7576 * Minor changes. 2026-08-12 11:20:40 +02:00
Antonio Cordero Balcazar 5abce373e0 * Updated model support.
* Fixed some model identifiers that had uppercase letters.
* Fixed missing env_info refactoring in one node.
* Improved documentation.
2026-08-07 14:24:08 +02:00
Antonio Cordero Balcazar 9692031cb5 * version bump 2026-08-05 14:27:26 +02:00
Antonio Cordero Balcazar 0e31e1cb33 * Refactored env_info into a dataclass. 2026-08-05 14:10:38 +02:00
Antonio Cordero Balcazar fd7c5596f0 * Converted do_combinatorial property to run_mode, with a new multiple mode. This changes the node properties in ComfyUI.
* Added option to set the default choice sampler.
* ComfyUI: Added ACBPPPRunModeOptions node to set run mode options.
* A1111: ppp object is now kept between executions, so cyclical state is saved.
* Added documentation in the cookbook regarding seed behavior and cyclical sampler resets.
* Adjusted the testing methods for more flexibility.
* Some refactoring.
2026-08-05 13:00:04 +02:00
Antonio Cordero Balcazar 6349038ce3 * Additional test for invalid wildcard detection. 2026-07-24 12:52:47 +02:00
Antonio Cordero Balcazar 90b4e2a791 * Added support for a fixed random sampler across combinations in combinatorial mode. 2026-07-21 19:39:25 +02:00
Antonio Cordero Balcazar 054598da0a * fix typo 2026-07-21 13:25:07 +02:00
Antonio Cordero Balcazar b9afef3205 * Random sampler stops combinatorial choices in combinatorial mode. 2026-07-20 17:40:54 +02:00
Antonio Cordero Balcazar 2eaf2b05fe * Improved detection of model class and some new models.
* Fixed seed size in ComfyUI to account for frontend limits.
* Better wildcard key validation.
* Some refactoring.
2026-07-20 15:45:04 +02:00
Antonio Cordero Balcazar 3a2571c008 * Support for adding more input vars from outside the prompt.
* Added input variables `_input_prev_pos_prompt` and `_input_prev_neg_prompt` for hiresfix phase of A1111 compatible hosts.
2026-06-12 22:53:21 +02:00
Antonio Cordero Balcazar 5e6460c099 * Added script tools/check_loras.py to check for problems with lora references in wildcards or mappings. 2026-06-12 16:21:03 +02:00
Antonio Cordero Balcazar d115ed2639 * Add "PPP Hires prompt" and "PPP Hires negative prompt" metadata to counteract A1111 hosts bug with bad "Hires prompt" / "Hires negative prompt" metadata when using "Batch count" > 1. 2026-06-12 14:05:24 +02:00
Antonio Cordero Balcazar dbe8a2b05d * Added option to have a default choice when no other choices are available after conditions (in choices and wildcards). 2026-06-08 16:49:15 +02:00
Antonio Cordero Balcazar 423faf1336 * Improved model class detection.
* ComfyUI: ignore errors on install script.
2026-06-08 13:07:20 +02:00
Antonio Cordero Balcazar 6802f328f6 * Fix pyproject to make the registry publishing action work again. 2026-06-05 11:31:34 +02:00
Antonio Cordero Balcazar 196978c074 * Pinning dependencies versions. 2026-06-04 18:19:59 +02:00
32 changed files with 1903 additions and 838 deletions
+24 -9
View File
@@ -65,7 +65,7 @@ def test_cl_combinatorial(self):
{} # expected variables (optional) {} # expected variables (optional)
), ),
], ],
combinatorial=True, ppp=self.init_ppp(None, run_mode=RUN_MODE.combinatorial),
) )
``` ```
@@ -79,7 +79,10 @@ def test_cl_combinatorial(self):
| `seed` | `int` | Optional, defaults to fixed seed | | `seed` | `int` | Optional, defaults to fixed seed |
| `ppp` | `PromptPostProcessor \| str \| None` | Supported values `"nocup"`, `"nostrict"` or a specific instance | | `ppp` | `PromptPostProcessor \| str \| None` | Supported values `"nocup"`, `"nostrict"` or a specific instance |
| `interrupted` | `bool` | Expected interrupt flag | | `interrupted` | `bool` | Expected interrupt flag |
| `combinatorial` | `bool` | Whether to run a combinatorial generation. If a specific ppp instance is used then it is ignored | | `specific_wc_folders` | `list[Path]` | Optional list of specific wildcard folders to use for this test |
| `specific_em_folders` | `list[Path]` | Optional list of specific extranetwork folders to use for this test |
| `input_vars` | `dict[str, Any]` | Optional dictionary of input variables to set before processing |
## Assertions ## Assertions
@@ -94,23 +97,35 @@ Do not use bare `assert` statements.
## Default Options & Environment ## Default Options & Environment
Override `self.defopts` or `self.def_env_info` to pass non-default options — do not hardcode option dicts from scratch: Override `self.defopts` or `self.def_env_info` to pass non-default options.
```python ```python
def test_cl_custom(self):
def test_cl_custom1(self): # only option changes
"""cleanup with custom separator"""
self.process(
InputTuple("a, , b", ""),
OutputTuple("a | b", ""),
ppp=self.init_ppp(
None,
keep_choices_order=True,
cup_do_cleanup=False,
run_mode=RUN_MODE.combinatorial,
),
)
def test_cl_custom2(self): # only environment changes
"""cleanup with custom separator""" """cleanup with custom separator"""
self.process( self.process(
InputTuple("a, , b", ""), InputTuple("a, , b", ""),
OutputTuple("a | b", ""), OutputTuple("a | b", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
self.def_env_info,
replace( replace(
self.defopts, self.def_env_info,
keep_choices_order=True, model_filename="./webui/models/Stable-diffusion/testmodel.safetensors",
cup_do_cleanup=False,
do_combinatorial=True,
), ),
self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
self.wildcards_obj, self.wildcards_obj,
+2
View File
@@ -8,3 +8,5 @@ logs
tests/tests_local.py tests/tests_local.py
tests/local_wildcards tests/local_wildcards
tests/logs tests/logs
tools/*.bat
dev/*.bat
+9 -2
View File
@@ -85,11 +85,18 @@ See the [cookbook](docs/COOKBOOK.md) for interesting usages.
## Tools ## Tools
A tool `tools/convert_styles.py` exists to convert A1111 or SD.Next styles into a wildcards file. ### User tools
* `tools/convert_styles.py`: converts A1111 or SD.Next styles into a wildcards file.
* `tools/check_loras.py`: checks for valid loras inside wildcards.
### Dev tools
* `dev/compare_models.py`: checks for new supported models in hosts.
## Contributing ## Contributing
To develop, I suggest doing so with the extension isolated from the UI (you can use a symlink to test it in the UI), and with its own virtual environment (venv or .venv), so the tests work and can be debugged properly. To develop, I suggest doing so with the extension isolated from the host UI (you can use a symlink to test it in the UI), and with its own virtual environment (venv or .venv), so the tests work and can be debugged properly.
## License ## License
+4
View File
@@ -10,8 +10,10 @@ from pathlib import Path
sys.path.append(str(Path(__file__).resolve().parent)) sys.path.append(str(Path(__file__).resolve().parent))
# pylint: disable=wrong-import-position
from .ppp_comfyui import ( from .ppp_comfyui import (
PromptPostProcessorComfyUINode, PromptPostProcessorComfyUINode,
PromptPostProcessorRunModeOptionsComfyUINode,
PromptPostProcessorWildcardOptionsComfyUINode, PromptPostProcessorWildcardOptionsComfyUINode,
PromptPostProcessorENMappingOptionsComfyUINode, PromptPostProcessorENMappingOptionsComfyUINode,
PromptPostProcessorSTNOptionsComfyUINode, PromptPostProcessorSTNOptionsComfyUINode,
@@ -22,6 +24,7 @@ from .ppp_comfyui import (
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"ACBPromptPostProcessor": PromptPostProcessorComfyUINode, "ACBPromptPostProcessor": PromptPostProcessorComfyUINode,
"ACBPPPRunModeOptions": PromptPostProcessorRunModeOptionsComfyUINode,
"ACBPPPWildcardOptions": PromptPostProcessorWildcardOptionsComfyUINode, "ACBPPPWildcardOptions": PromptPostProcessorWildcardOptionsComfyUINode,
"ACBPPPENMappingOptions": PromptPostProcessorENMappingOptionsComfyUINode, "ACBPPPENMappingOptions": PromptPostProcessorENMappingOptionsComfyUINode,
"ACBPPPSendToNegativeOptions": PromptPostProcessorSTNOptionsComfyUINode, "ACBPPPSendToNegativeOptions": PromptPostProcessorSTNOptionsComfyUINode,
@@ -31,6 +34,7 @@ NODE_CLASS_MAPPINGS = {
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"ACBPromptPostProcessor": "ACB Prompt Post Processor", "ACBPromptPostProcessor": "ACB Prompt Post Processor",
"ACBPPPRunModeOptions": "ACB PPP Run Mode Options",
"ACBPPPWildcardOptions": "ACB PPP Wildcard Options", "ACBPPPWildcardOptions": "ACB PPP Wildcard Options",
"ACBPPPENMappingOptions": "ACB PPP ExtraNetwork Mapping Options", "ACBPPPENMappingOptions": "ACB PPP ExtraNetwork Mapping Options",
"ACBPPPSendToNegativeOptions": "ACB PPP Send-To-Negative Options", "ACBPPPSendToNegativeOptions": "ACB PPP Send-To-Negative Options",
+48 -9
View File
@@ -12,7 +12,7 @@ The model variants now support regular expressions instead of a list of strings
## Important ## Important
**Beware of the combinatorial mode with no limits**. Even very few choice/wildcard constructs can cause a *combinatorial explosion*! **Beware of the combinatorial mode with no limits**. Even very few choice/wildcard constructs can cause a *combinatorial explosion*! In ComfyUI there is no other limit, but in the other hosts this is also limited by the batch count/size.
The console log can help you determine the number of combinations that it is trying to generate. There will be an **"Estimated combinations"** message that shows an estimate. You can try first with a limit of 1, then check this message in the log. But note that it is a lower bound estimate, and there could be more combinations. The console log can help you determine the number of combinations that it is trying to generate. There will be an **"Estimated combinations"** message that shows an estimate. You can try first with a limit of 1, then check this message in the log. But note that it is a lower bound estimate, and there could be more combinations.
@@ -35,14 +35,13 @@ Inputs:
* **process_wildcards**: Activates the wildcard processing. * **process_wildcards**: Activates the wildcard processing.
* **do_cleanup**: Activates the cleanup processing. * **do_cleanup**: Activates the cleanup processing.
* **cleanup_variables**: Do a cleanup of the output variables (depends on do_cleanup). * **cleanup_variables**: Do a cleanup of the output variables (depends on do_cleanup).
* **do_combinatorial**: Activates combinatorial mode, where the output are all the combinations of choices/wildcards of the prompt. * **results_file**: Filename to save processing results. Supports `%datetime%`, `%date%`, `%time%`, and `%host%` tokens. The file extension determines the format: `.yaml`/`.yml`, `.jsonl`, `.csv`, or plain text for any other extension. Relative paths are resolved against the extension's `logs` folder. Leave empty to disable.
* **combinatorial_shuffle**: It shuffles the combinatorial results. * **run_mode**: Sets how the process works. `single` or `multiple` for regular one or more results, or `combinatorial` for combinatorial mode, where the output are all the combinations of choices/wildcards of the prompt.
* **combinatorial_limit**: Limit for the number of generated combinations.
* **wc_options**: Connection to a Wildcards options node. * **wc_options**: Connection to a Wildcards options node.
* **stn_options**: Connection to a Send-To-Negative options node. * **stn_options**: Connection to a Send-To-Negative options node.
* **cup_options**: Connection to a Cleanup options node. * **cup_options**: Connection to a Cleanup options node.
* **en_options**: Connection to a ExtraNetworkMapping options node. * **en_options**: Connection to a ExtraNetworkMapping options node.
* **results_file**: Filename to save processing results. Supports `%datetime%`, `%date%`, `%time%`, and `%host%` tokens. The file extension determines the format: `.yaml`/`.yml`, `.jsonl`, `.csv`, or plain text for any other extension. Relative paths are resolved against the extension's `logs` folder. Leave empty to disable. * **rm_options**: Connection to a Run Mode options node.
The options nodes are optional. If you don't need to change any of the default values then you don't need to use them. The options nodes are optional. If you don't need to change any of the default values then you don't need to use them.
@@ -58,7 +57,19 @@ Outputs:
* **neg_prompt**: Resulting negative prompt. * **neg_prompt**: Resulting negative prompt.
* **variables**: Resulting output variables. * **variables**: Resulting output variables.
The outputs are lists, and in combinatorial mode there will be multiple elements that *ComfyUI* will process sequentially. The outputs are lists, and in combinatorial/multiple modes there will be multiple elements that *ComfyUI* will process sequentially.
The run_mode can be explained like this:
* **single**: only one result is returned.
* **multiple**: multiple results are returned, the count is in results_limit.
* **combinatorial**: all the combinations are returned, up to results_limit.
In `single` and `multiple` modes the default choice sampler is the one set in `default_sampler`. In `combinatorial` mode the default sampling is equivalent to `cyclical`. In all modes specified samplers are respected. The value of random samplers in `combinatorial` mode depends on `comb_random_fixed`.
`Multiple` mode with a default of `cyclical` sampler is very similar to `combinatorial`. The only difference is that `comb_ramdom_fixed` does not apply and random samplers are thus not cached.
Single mode with a default of cyclical sampler can be used as similar to combinatorial but in separated *ComfyUI* runs instead of one.
### ACB PPP Select Variable node ### ACB PPP Select Variable node
@@ -94,6 +105,20 @@ Output:
* **prompt**: concatenated result. * **prompt**: concatenated result.
### ACB PPP Run Mode Options node
Options for the run mode, in case you want to change them from the defaults.
* **results_limit**: Limit for the number of generated results (except in `single` mode). Important for combinatorial mode.
* **results_shuffle**: It shuffles the results.
* **comb_random_fixed**: If True all specified random samplers will have a fixed value across the combinations.
* **default_sampler**: The default choice sampler when not specified (in non combinatorial mode). Also applies to extranetwork mapping selection.
* **next_seed**: Choose what to do with the seed in the following prompts in `multiple` or `combinatorial` mode. Value can be:
* `randomize`: Next prompts will have a random seed (default).
* `input`: Next prompts will have the same seed as the input (or first prompt).
* `increment`: The seeds in next prompts will increment.
* `decrement`: The seeds in next prompts will decrement.
### ACB PPP Wildcard Options node ### ACB PPP Wildcard Options node
Options for wildcard processing, in case you want to change them from the defaults. Options for wildcard processing, in case you want to change them from the defaults.
@@ -149,9 +174,23 @@ Options for extranetworks mapping, in case you want to change them from the defa
* **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. * **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. * **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. * **Incremental seed**: When using a batch you can use this to set the rest of the prompt seeds with consecutive values.
* **Combinatorial mode**: Generate all possible prompt combinations (from choices and wildcards) and cycle through them to fill the batch. * **Run mode**: Sets how the process works. `single` or `multiple` for regular one or more results, or `combinatorial` for combinatorial mode, where the output are all the combinations of choices/wildcards of the prompt.
* **Shuffle combinations**: It shuffles the combinatorial results. * **Results limit**: Maximum number of combinations to generate (0 = no limit). The actual maximum limit is the number of images (batch size * count).
* **Combinations limit**: Maximum number of combinations to generate (0 = no limit). The actual maximum limit is the number of images (batch size * count). * **Shuffle results**: It shuffles the results.
* **Fix random sampler across combinations**: If checked all specified random samplers will have a fixed value across the combinations.
* **Default sampler**: The default choice sampler when not specified.
The `Run mode` can be explained like this:
* **single**: only one result is returned (for each image).
* **multiple**: multiple results are returned, the count is in `Results limit` (but limited by the number of images). The results are assigned to the images. If there are more images than results, the results wrap.
* **combinatorial**: all the combinations are returned, up to `Results limit` (but limited by the number of images). The results are assigned to the images. If there are more images than results, the results wrap.
In single and multiple modes the default choice sampler is the one set in `Default sampler`. In combinatorial mode the default sampling is equivalent to cyclical. In all modes specified samplers are respected. The value of random samplers in combinatorial mode depends on `Fix random sampler across combinations`.
Multiple mode with a default of cyclical sampler is very similar to combinatorial. The only difference is that `Fix random sampler across combinations` does not apply and random samplers are thus not cached. There is no reason to use this mode with a random default sampler, since it would be the same as in single mode.
Single mode with a default of cyclical sampler can be used as similar to combinatorial but in separated runs instead of one.
### General settings ### General settings
+52 -6
View File
@@ -496,6 +496,52 @@ You can leave the model and modelname inputs disconnected/empty and set the `_mo
You can also set and extract user variables for other ksampler inputs, like sampler, scheduler, steps, cfg and latent size. You can also set and extract user variables for other ksampler inputs, like sampler, scheduler, steps, cfg and latent size.
## Seed behavior
The seed determines which choices are picked for wildcards/choices with `~` (random) samplers and other random selections.
This seed can be different than the one used for the image.
Each host handles it differently.
### A1111 and derivatives
By default, PPP uses the same seed for image and prompt, which A1111 auto-increments across the batch. Each image therefore gets independently seeded wildcard expansions.
The seeds are pre-calculated before the prompt postprocess begins.
The extension provides options to change this:
- **Force equal seeds**: sets every seed in the batch to the first one before processing. All images get the same expansion.
- **Unlink seed**: separates the prompt seed from the image seed. The table below shows how it behaves depending on the seed value and the *Incremental seed* toggle:
| Seed value | Incremental | Prompt seed per image |
|------------|-------------|------------------------------------------|
| -1 | Yes | Random base seed, then base+1, base+2, … |
| -1 | No | Independent random seed per image |
| N | Yes | N, N+1, N+2, … |
| N | No | N for every image |
If a subseed strength is set, the effective seed is `subseed × strength + seed × (1 − strength)` per image.
In **multiple** or **combinatorial** run mode, results are generated using their corresponding seed in the batch, both for prompts and images. The `next_seed` option is set to `input` with these hosts (and a list of them is provided) because we need the image seeds pre-calculated.
### ComfyUI
The seed is an explicit node input (default: -1 for random, which gets an actual value as soon as possible and is considered the starting seed). In multiple result modes (`multiple` and `combinatorial`) the `next_seed` option is used to calculate the rest of the seeds.
Each result seed is added as the output variable `_output_seed`, so it can be extracted and used as the image seed.
## When the cyclical sampler resets
The `@` (cyclical) sampler tracks its position so each call advances through combinations in order.
The state is retained across executions within the same session. Each execution with the same prompts advances the position. The cycle only resets when either the positive or negative prompt text changes.
## Keeping references up to date
After updating your loras (deleting old ones, updating to new versions) run the `tools\check_loras.py` script to check if there are broken references in your wildcards or extranetwork mappings.
## Debugging tips ## Debugging tips
When something isn't generating as expected, the debug setting is your first tool. Enable it in the extension settings; it will log all system variables at generation time, which tells you exactly what values are available for your conditions. When something isn't generating as expected, the debug setting is your first tool. Enable it in the extension settings; it will log all system variables at generation time, which tells you exactly what values are available for your conditions.
@@ -550,12 +596,12 @@ Common use cases:
Set the `results_file` option (in the extension settings for *A1111*, or the `results_file` input on the main node for *ComfyUI*) to a filename. The extension determines the output format from the file extension: Set the `results_file` option (in the extension settings for *A1111*, or the `results_file` input on the main node for *ComfyUI*) to a filename. The extension determines the output format from the file extension:
| Extension | Format | | Extension | Format |
|------------------|----------------------------------------------------| |------------------|---------------------------------------------|
| `.yaml` / `.yml` | YAML list of records | | `.yaml` / `.yml` | YAML list of records |
| `.jsonl` | JSON Lines, one JSON object per line | | `.jsonl` | JSON Lines, one JSON object per line |
| `.csv` | CSV with a header row (semicolon-delimited) | | `.csv` | CSV with a header row (semicolon-delimited) |
| anything else | Plain text with labelled sections | | anything else | Plain text with labelled sections |
Each record contains five sections: `options` (the PPP settings that were active), `inputs` (seed, prompts), `system` (system variables like `_modelclass`), `results` (the final positive and negative prompts), and `variables` (any user variables that were set). Each record contains five sections: `options` (the PPP settings that were active), `inputs` (seed, prompts), `system` (system variables like `_modelclass`), `results` (the final positive and negative prompts), and `variables` (any user variables that were set).
+53 -41
View File
@@ -1,5 +1,9 @@
# Prompt PostProcessor syntax # 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 ## 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). 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).
@@ -37,7 +41,7 @@ There is also a format where instead of `parameters$$` you just put the sampler,
The construct parameters can be written with the following options (all are optional): The construct parameters can be written with the following options (all are optional):
* `~` (random) or `@` (cyclical): sampler (for compatibility with *Dynamic Prompts*). The cyclical sampler cycles through all combinations in order across consecutive `process_prompt` calls, resuming where the previous call left off (as long as the input prompt and negative prompt do not change). * `~` (random) or `@` (cyclical): sampler (for compatibility with *Dynamic Prompts*). The cyclical sampler cycles through all combinations in order across consecutive `process_prompt` calls, resuming where the previous call left off (as long as the input prompt and negative prompt do not change). In combinatorial mode you can use the random sampler to stop a specific choice/wildcard from being expanded into all its combinations.
* `r`: means it allows repetition of the choices. * `r`: means it allows repetition of the choices.
* `o`: means it is optional, and no error will be raised if there are no choices to select from. * `o`: means it is optional, and no error will be raised if there are no choices to select from.
* `n` or `n-m` or `n-` or `-m`: number or range of choices to select. Allows zero as the start of a range. Default is 1. * `n` or `n-m` or `n-` or `-m`: number or range of choices to select. Allows zero as the start of a range. Default is 1.
@@ -53,6 +57,7 @@ The choice options are as follows:
* `'identifiers'`: comma separated labels for the choice (optional, quotes can be single or double). Only makes sense inside a wildcard definition. Can be used when specifying the wildcard to select this specific choice. It's case insensitive. * `'identifiers'`: comma separated labels for the choice (optional, quotes can be single or double). Only makes sense inside a wildcard definition. Can be used when specifying the wildcard to select this specific choice. It's case insensitive.
* `n`: weight of the choice (optional, default 1). * `n`: weight of the choice (optional, default 1).
* `if condition`: filters out the choice if the condition is false (optional; this is an extension to the *Dynamic Prompts* syntax). Same conditions as in the `if` command. * `if condition`: filters out the choice if the condition is false (optional; this is an extension to the *Dynamic Prompts* syntax). Same conditions as in the `if` command.
* `else`: flag to indicate this is the choice to use if no other choice is available after conditions. It won't be considered if other choices are available. Also optional.
* `::`: end of choice options (not optional if any options) * `::`: end of choice options (not optional if any options)
Whitespace is allowed between parameters/options. Whitespace is allowed between parameters/options.
@@ -62,7 +67,7 @@ The only command available is `include wildcard`, which will include the choices
These are examples of formats you can use to insert a choice construct: These are examples of formats you can use to insert a choice construct:
| Construct | Result | | Construct | Result |
| --------- | ------ | |---------------------------------------------------|---------------------------------------------------------------------------------------------------------------------------------------|
| `{choice1\|5::choice2\|3::choice3}` | select 1 choice, two of them have weights | | `{choice1\|5::choice2\|3::choice3}` | select 1 choice, two of them have weights |
| `{3$$choice1\|5 if _is_sd1::choice2\|choice3}` | select 3 choices, one has a weight and a condition | | `{3$$choice1\|5 if _is_sd1::choice2\|choice3}` | select 3 choices, one has a weight and a condition |
| `{2-3$$2::choice1\|choice2\|choice3}` | select 2 to 3 choices, one of them has a weight | | `{2-3$$2::choice1\|choice2\|choice3}` | select 2 to 3 choices, one of them has a weight |
@@ -111,7 +116,7 @@ The variable value only applies during the evaluation of the selected choices an
These are examples of formats you can use to insert a wildcard: These are examples of formats you can use to insert a wildcard:
| Construct | Result | | Construct | Result |
| --------- | ------ | |--------------------------------------|--------------------------------------------------------------------------|
| `__wildcard__` | select 1 choice | | `__wildcard__` | select 1 choice |
| `__path/wildcard'0'__` | select the first choice | | `__path/wildcard'0'__` | select the first choice |
| `__path/wildcard'1-2'__` | select the second or third choice | | `__path/wildcard'1-2'__` | select the second or third choice |
@@ -142,6 +147,7 @@ The choices of the wildcard follow the same format as in the choices construct,
If using the object format for a choice you can use the following in addition to the standard `weight` and `text`/`content`: If using the object format for a choice you can use the following in addition to the standard `weight` and `text`/`content`:
* `if`: the condition (a string) * `if`: the condition (a string)
* `else`: flag (boolean) to indicate this is the choice to use if no other choice is available after conditions. It won't be considered if other choices are available.
* `labels`: list of labels (an array of strings) * `labels`: list of labels (an array of strings)
* `command`: indicates the content is a command (a boolean) * `command`: indicates the content is a command (a boolean)
@@ -192,7 +198,7 @@ This command can be used to set a default filter for a wildcard, before it is us
The format is: The format is:
| Construct | Meaning | | Construct | Meaning |
| --------- | ------- | |-----------------------------------------------|--------------------|
| `<ppp:setwcdeffilter 'identifier' 'filter'/>` | Sets a filter | | `<ppp:setwcdeffilter 'identifier' 'filter'/>` | Sets a filter |
| `<ppp:setwcdeffilter 'identifier'/>` | Removes the filter | | `<ppp:setwcdeffilter 'identifier'/>` | Removes the filter |
@@ -210,28 +216,34 @@ All these variables can be used to output content or behave differently based on
Names starting with an underscore are reserved for system variables: Names starting with an underscore are reserved for system variables:
| System variable | Value | | System variable | Value |
| --------------- | ----- | |--------------------------|---------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| `_model` | the model identifier (`sd1`, `sd2`, `sdxl`, `sd3`, `flux`, `auraflow`). `_sd` also works but is deprecated. | | `_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. | | `_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). | | `_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. | | `_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_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_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_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_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_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_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_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. | | `_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. | | `_opt_...` | All the options. |
| `_input_seed` | The seed used. | | `_input_seed` | The starting seed used. |
| `_input_pos_prompt` | The original positive prompt. | | `_input_pos_prompt` | The original positive prompt. |
| `_input_neg_prompt` | The original negative 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. |
| `_output_seed` | The seed used for the specific result. This variable changes for each result depending on the `next_seed` option. You can feed it to the ksampler. |
> [!NOTE] > [!NOTE]
> The model path is relative to the checkpoint/difussion_models folder, just as it appears in the load nodes. > 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 ## Set command
This command sets the value of a variable that can be checked later. This command sets the value of a variable that can be checked later.
@@ -249,14 +261,14 @@ The `add` and `ifundefined` modifiers are mutually exclusive and cannot be used
The *Dynamic Prompts* format also works: The *Dynamic Prompts* format also works:
| Construct | Meaning | | Construct | Meaning |
| --------- | ------- | |-----------------|----------------------|
| `${var=value}` | regular evaluation | | `${var=value}` | regular evaluation |
| `${var=!value}` | immediate evaluation | | `${var=!value}` | immediate evaluation |
If also supports the addition and undefined check as an extension of the *Dynamic Prompts* format: If also supports the addition and undefined check as an extension of the *Dynamic Prompts* format:
| Construct | Meaning | | Construct | Meaning |
| --------- | ------- | |------------------|--------------------------------------|
| `${var+=value}` | equivalent to `add` | | `${var+=value}` | equivalent to `add` |
| `${var+=!value}` | equivalent to `evaluate add` | | `${var+=!value}` | equivalent to `evaluate add` |
| `${var?=value}` | equivalent to `ifundefined` | | `${var?=value}` | equivalent to `ifundefined` |
@@ -290,14 +302,14 @@ This command prints the value of a variable, or the specified default if it does
The format is: The format is:
| Construct | | Construct |
| --------- | |----------------------------------------|
| `<ppp:echo varname/>` | | `<ppp:echo varname/>` |
| `<ppp:echo varname>default<ppp:/echo>` | | `<ppp:echo varname>default<ppp:/echo>` |
The *Dynamic Prompts* format is: The *Dynamic Prompts* format is:
| Construct | | Construct |
| --------- | |----------------------|
| `${varname}` | | `${varname}` |
| `${varname:default}` | | `${varname:default}` |
@@ -310,7 +322,7 @@ There is support for array variables. They use brackets `[]` to differenciate fr
They can be initialized in several ways: They can be initialized in several ways:
| Construct | Meaning | | Construct | Meaning |
| --------- | ------- | |----------------------------|------------------------------------------------------------------|
| `${var[]=value}` | initialize and set the first value | | `${var[]=value}` | initialize and set the first value |
| `${var[]=*()}` | initialize an empty array | | `${var[]=*()}` | initialize an empty array |
| `${var[]=*var2[]}` | initialize an array from another array | | `${var[]=*var2[]}` | initialize an array from another array |
@@ -330,13 +342,13 @@ And they can be accesed/echoed with:
* A hash inside the brackets is used to get the length of the array * A hash inside the brackets is used to get the length of the array
* An ampersand followed by a string inside the brackets (with quotes) is used to get the full array joined with a separator. * An ampersand followed by a string inside the brackets (with quotes) is used to get the full array joined with a separator.
| Construct | Meaning | | Construct | Meaning |
| --------- | ------- | |---------------------|---------------------------------------------|
| `${var[]}` | echo all elements with a default separator | | `${var[]}` | echo all elements with a default separator |
| `${var[&' / ']}` | echo all elements with a specific separator | | `${var[&' / ']}` | echo all elements with a specific separator |
| `${var[n]}` | echo an element from the array | | `${var[n]}` | echo an element from the array |
| `${var[n]:default}` | echo an element with a default | | `${var[n]:default}` | echo an element with a default |
| `${var[#]}` | echo the length of the array | | `${var[#]}` | echo the length of the array |
## If command ## If command
@@ -351,7 +363,7 @@ Any `elif`s (there can be multiple) and the `else` are optional.
The `conditionN` is a boolean expression, which can use `and`, `or`, `not` and grouping with parentheses, and where the simplest expression can be: The `conditionN` is a boolean expression, which can use `and`, `or`, `not` and grouping with parentheses, and where the simplest expression can be:
| Construct | Meaning | | Construct | Meaning |
| --------- | ------- | |-------------------------------------|---------------------------------------------------------------------------------|
| `operand` | check truthyness of the operand, meaning not zero, empty string nor empty array | | `operand` | check truthyness of the operand, meaning not zero, empty string nor empty array |
| `operand1 [not] operation operand2` | compare the operands | | `operand1 [not] operation operand2` | compare the operands |
@@ -368,7 +380,7 @@ The supported operations are: `eq`, `ne`, `gt`, `lt`, `ge`, `le`, `in`, `any_in`
This list shows what they do depending on the kind of operand (R = regular variable, A = array variable). This list shows what they do depending on the kind of operand (R = regular variable, A = array variable).
| Operation | R1 op R2 | A1 op A2 | A1 op R2 | R1 op A2 | | Operation | R1 op R2 | A1 op A2 | A1 op R2 | R1 op A2 |
| --------- | -------- | -------- | -------- | -------- | |----------------|--------------------------|-------------------|------------------------------------------------|------------------------------------------------|
| `eq` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | | `eq` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise |
| `ne` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | | `ne` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise |
| `gt` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise | | `gt` | OK | OK (pairwise) | Error in strict mode, all A1 with R2 otherwise | Error in strict mode, R1 with all A2 otherwise |
@@ -458,10 +470,10 @@ extnettype:
Used like this: Used like this:
| Construct | Meaning | | Construct | Meaning |
| --------- | ------- | |--------------------------------------------------------|-------------------------------------|
| `<ppp:ext $lora mappingname/>` | Mapping without additional triggers | | `<ppp:ext $lora mappingname/>` | Mapping without additional triggers |
| `<ppp:ext $lora mappingname>inline triggers<ppp:/ext>` | Mapping with additional triggers | | `<ppp:ext $lora mappingname>inline triggers<ppp:/ext>` | Mapping with additional triggers |
Each mapping can have any number of elements in its list of mappings. There are no mandatory properties for a mapping. The properties mean the following: Each mapping can have any number of elements in its list of mappings. There are no mandatory properties for a mapping. The properties mean the following:
@@ -480,7 +492,7 @@ See the file in the tests folder as an example.
The new format for this command is like this: The new format for this command is like this:
| Construct | Meaning | | Construct | Meaning |
| --------- | ------- | |---------------------------------------|--------------------------------------------------------------------------------------|
| `<ppp:stn position>content<ppp:/stn>` | send to negative prompt | | `<ppp:stn position>content<ppp:/stn>` | send to negative prompt |
| `<ppp:stn iN/>` | insertion point to be used in the negative prompt as destination for the pN position | | `<ppp:stn iN/>` | insertion point to be used in the negative prompt as destination for the pN position |
+5 -4
View File
@@ -165,27 +165,28 @@ extranetworktag: "<" /(?!ppp:)\w+:/ encontent ">"
choicesoptions_sep: "$$" plain choicesoptions_sep: "$$" plain
// choice options // choice options
choice: [ [ _WHITESPACE? choiceiscmd ] [ _WHITESPACE? choicelabels ] [ _WHITESPACE? choiceweight ] [ _WHITESPACE? choiceif ] _WHITESPACE? "::" ] choicevalue choice: [ [ _WHITESPACE? choiceiscmd ] [ _WHITESPACE? choicelabels ] [ _WHITESPACE? choiceweight ] [ _WHITESPACE? ( choiceif | choiceelse ) ] _WHITESPACE? "::" ] choicevalue
choiceiscmd: /%/ // the option text is a special command choiceiscmd: /%/ // the option text is a special command
choicelabels: ( /"/ IDENTIFIER ( _WHITESPACE? "," _WHITESPACE? IDENTIFIER )* /"/ ) choicelabels: ( /"/ IDENTIFIER ( _WHITESPACE? "," _WHITESPACE? IDENTIFIER )* /"/ )
| ( /'/ IDENTIFIER ( _WHITESPACE? "," _WHITESPACE? IDENTIFIER )* /'/ ) | ( /'/ IDENTIFIER ( _WHITESPACE? "," _WHITESPACE? IDENTIFIER )* /'/ )
choiceweight: NUMBER choiceweight: NUMBER
choiceif: "if" _WHITESPACE condition choiceif: "if" _WHITESPACE condition
choiceelse: "else"
choicevalue: content_choice choicevalue: content_choice
//#endif //#endif
//#if ALLOW_CHOICES //#if ALLOW_CHOICES
// choices construct // choices construct
choices: "{" [ choicesoptions_sampler | ( choicesoptions _WHITESPACE? "$$" ) ] choice ( "|" choice )* "}" choices: "{" [ choicesoptions_sampler | ( choicesoptions "$$" ) ] choice ( "|" choice )* "}"
//#endif //#endif
//#if ALLOW_WILDCARDS //#if ALLOW_WILDCARDS
// wildcard definition options // wildcard definition options
wcdefoptions: [ choicesoptions_sampler ] [ _WHITESPACE? choicesoptions_flags ] [ _WHITESPACE? choicesoptions_range ] [ _WHITESPACE? wcdescription ] [ _WHITESPACE? choicesoptions_sep ] wcdefoptions: [ choicesoptions_sampler ] [ _WHITESPACE? choicesoptions_flags ] [ _WHITESPACE? choicesoptions_range ] [ _WHITESPACE? wcdescription ] [ _WHITESPACE? choicesoptions_sep ] "$$"
wcdescription: STRING wcdescription: STRING
// wildcards construct // wildcards construct
wildcard: "__" [ choicesoptions_sampler | ( choicesoptions _WHITESPACE? "$$" ) ] wildcard_name [ wc_filter ] [ wildcardvar ] "__" wildcard: "__" [ choicesoptions_sampler | ( choicesoptions "$$" ) ] wildcard_name [ wc_filter ] [ wildcardvar ] "__"
wc_filter_nums: IDENTIFIER | INDEX | (INDEX /-/ INDEX) wc_filter_nums: IDENTIFIER | INDEX | (INDEX /-/ INDEX)
//#if ALLOW_COMMVARS //#if ALLOW_COMMVARS
wildcard_name.2: ( WC_NAME_PLAIN_START | variableuse | commandecho ) ( WC_NAME_PLAIN | variableuse | commandecho )* wildcard_name.2: ( WC_NAME_PLAIN_START | variableuse | commandecho ) ( WC_NAME_PLAIN | variableuse | commandecho )*
+5
View File
@@ -1,3 +1,6 @@
"""
Install dependencies for Prompt Post-Processor extension. For A1111 hosts.
"""
from pathlib import Path from pathlib import Path
requirements_filename = str(Path(__file__).resolve().parent / "requirements.txt") requirements_filename = str(Path(__file__).resolve().parent / "requirements.txt")
@@ -9,3 +12,5 @@ try:
except ImportError: except ImportError:
import launch import launch
launch.run_pip(f'install -r "{requirements_filename}"', "requirements for Prompt Post-Processor") launch.run_pip(f'install -r "{requirements_filename}"', "requirements for Prompt Post-Processor")
except Exception as e:
pass
+142 -66
View File
@@ -20,6 +20,9 @@ from ppp_classes import (
HostConfig, HostConfig,
ModelConfig, ModelConfig,
ModelDetectConfig, ModelDetectConfig,
PPPEnvInfo,
PPPException,
RUN_MODE,
VariantConfig, VariantConfig,
PPPConfig, PPPConfig,
IFWILDCARDS_CHOICES, IFWILDCARDS_CHOICES,
@@ -33,7 +36,15 @@ from ppp_variables import VariableRepository, VariableEntry, VariableValue
from ppp_logging import DEBUG_LEVEL, log from ppp_logging import DEBUG_LEVEL, log
from ppp_tree import TreeProcessor from ppp_tree import TreeProcessor
from ppp_utils import escape_single_quotes, get_version_from_pyproject from ppp_utils import escape_single_quotes, get_version_from_pyproject
from ppp_common import get_model_class_from_filename, load_grammar, parse_prompt, preprocess_grammar, warn_or_stop from ppp_common import (
WARN_STOP_WHERE,
clamp_host_bits,
get_model_class_from_filename,
load_grammar,
parse_prompt,
preprocess_grammar,
warn_or_stop,
)
from ppp_wildcards import PPPWildcards from ppp_wildcards import PPPWildcards
from ppp_enmappings import PPPExtraNetworkMappings from ppp_enmappings import PPPExtraNetworkMappings
@@ -70,10 +81,13 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
DEFAULT_CUP_MERGE_ATTENTION = defopt["cup_merge_attention"] DEFAULT_CUP_MERGE_ATTENTION = defopt["cup_merge_attention"]
DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = defopt["cup_remove_extranetwork_tags"] DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS = defopt["cup_remove_extranetwork_tags"]
DEFAULT_STRICT_OPERATORS = defopt["strict_operators"] DEFAULT_STRICT_OPERATORS = defopt["strict_operators"]
DEFAULT_DO_COMBINATORIAL = defopt["do_combinatorial"] DEFAULT_RUN_MODE = defopt["run_mode"].value
DEFAULT_COMBINATORIAL_SHUFFLE = defopt["combinatorial_shuffle"] DEFAULT_RESULTS_SHUFFLE = defopt["results_shuffle"]
DEFAULT_COMBINATORIAL_LIMIT = defopt["combinatorial_limit"] DEFAULT_RESULTS_LIMIT = defopt["results_limit"]
DEFAULT_COMB_RANDOM_FIXED = defopt["comb_random_fixed"]
DEFAULT_DEFAULT_SAMPLER = defopt["default_sampler"].value
DEFAULT_RESULTS_FILE = defopt["results_file"] DEFAULT_RESULTS_FILE = defopt["results_file"]
DEFAULT_NEXT_SEED = defopt["next_seed"].value
WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK ' WILDCARD_WARNING = '(WARNING TEXT "INVALID WILDCARD" IN BRIGHT RED:1.5)\nBREAK '
WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK " WILDCARD_STOP = "INVALID WILDCARD! {0}\nBREAK "
@@ -82,7 +96,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
def __init__( def __init__(
self, self,
logger: logging.Logger, logger: logging.Logger,
env_info: dict[str, Any], env_info: PPPEnvInfo,
options: PPPStateOptions, options: PPPStateOptions,
grammar_content: Optional[str] = None, grammar_content: Optional[str] = None,
interrupt: Optional[Callable] = None, interrupt: Optional[Callable] = None,
@@ -95,7 +109,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
Args: Args:
logger: The logger object. logger: The logger object.
interrupt: The interrupt function. interrupt: The interrupt function.
env_info: A dictionary with information for the environment and loaded model. env_info: Environment and model information.
options: The options object for configuring PPP behavior. options: The options object for configuring PPP behavior.
grammar_content: Optional. The grammar content to be used for parsing. grammar_content: Optional. The grammar content to be used for parsing.
wildcards_obj: Optional. The wildcards object to be used for processing wildcards. wildcards_obj: Optional. The wildcards object to be used for processing wildcards.
@@ -272,7 +286,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None, exc_info=None): def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None, exc_info=None):
log(self.logger, self.debug_level, kind, message, min_level, exc_info=exc_info) log(self.logger, self.debug_level, kind, message, min_level, exc_info=exc_info)
def __load_config_and_detect(self, env_info: dict[str, Any]) -> HostConfig: def __load_config_and_detect(self, env_info: PPPEnvInfo) -> HostConfig:
"""Loads config files, performs model detection, and returns the resolved host config.""" """Loads config files, performs model detection, and returns the resolved host config."""
main_folder = Path(__file__).resolve().parent main_folder = Path(__file__).resolve().parent
default_config_file = str(main_folder / "ppp_config.yaml.defaults") default_config_file = str(main_folder / "ppp_config.yaml.defaults")
@@ -292,15 +306,15 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
raise PPPInterrupt(errmsg) raise PPPInterrupt(errmsg)
self.log(logging.WARNING, errmsg) self.log(logging.WARNING, errmsg)
app = env_info.get("app", "") app = env_info.app
user_config_file = env_info.get("ppp_config", "") user_config_file = env_info.ppp_config or ""
if isinstance(user_config_file, dict): if isinstance(user_config_file, dict):
user_cfg, _ = self.__parse_configuration(user_config_file, "forced configuration") user_cfg, _ = self.__parse_configuration(user_config_file, "forced configuration")
else: else:
if user_config_file == "": if user_config_file == "":
if app == SUPPORTED_APPS.comfyui.value: if app == SUPPORTED_APPS.comfyui:
try: try:
import folder_paths # type: ignore import folder_paths # type: ignore # pylint: disable=import-outside-toplevel,import-error
user_dir = folder_paths.get_user_directory() user_dir = folder_paths.get_user_directory()
if user_dir and Path(user_dir).is_dir(): if user_dir and Path(user_dir).is_dir():
@@ -362,7 +376,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
self.known_models: list[str] = list(self.models_config.keys()) self.known_models: list[str] = list(self.models_config.keys())
# Patch for tests (copy comfyui) # Patch for tests (copy comfyui)
if app == "tests": if app == SUPPORTED_APPS.tests:
if self.config.hosts is None: if self.config.hosts is None:
self.config.hosts = {} self.config.hosts = {}
self.config.hosts.setdefault("tests", HostConfig()) self.config.hosts.setdefault("tests", HostConfig())
@@ -373,10 +387,10 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
model.detect = {} model.detect = {}
model.detect.setdefault("tests", model.detect.get("comfyui", None)) model.detect.setdefault("tests", model.detect.get("comfyui", None))
host_config: HostConfig | None = (self.config.hosts or {}).get(app) host_config: HostConfig | None = (self.config.hosts or {}).get(app.value)
if host_config is None: if host_config is None:
raise PPPInterrupt( raise PPPInterrupt(
f"No host configuration found for app '{escape_single_quotes(app)}'. Please check your configuration." f"No host configuration found for app '{escape_single_quotes(app.value)}'. Please check your configuration."
) )
# Update env_info with model detection # Update env_info with model detection
@@ -394,41 +408,48 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
) )
self.log( self.log(
logging.DEBUG, logging.DEBUG,
f"Host configuration ({escape_single_quotes(app)}): {host_config}", f"Host configuration ({escape_single_quotes(app.value)}): {host_config}",
min_level=DEBUG_LEVEL.minimal, min_level=DEBUG_LEVEL.minimal,
) )
return host_config return host_config
def __run_model_detection(self, env_info: dict[str, Any]) -> None: def __run_model_detection(self, env_info: PPPEnvInfo) -> None:
"""Updates the is_* model detection flags in env_info based on model_class.""" """Updates the is_* model detection flags in env_info based on model_class."""
prop_base = env_info.get("property_base", None) prop_base = env_info.property_base
model_class = env_info.get("model_class", "") model_class = env_info.model_class
model_name = env_info.get("model_filename", "") model_name = env_info.model_filename
app = env_info.get("app", "") app = env_info.app
if not model_class and model_name and app == SUPPORTED_APPS.comfyui.value: if not model_class and model_name and app == SUPPORTED_APPS.comfyui:
model_class = get_model_class_from_filename(model_name) try:
model_class = get_model_class_from_filename(Path(model_name))
except PPPException as e:
self.log(
logging.WARNING,
f"Could not detect model class from filename '{model_name}': {e}",
min_level=DEBUG_LEVEL.minimal,
)
if model_class: if model_class:
env_info["model_class"] = model_class env_info.model_class = model_class
self.log( self.log(
logging.DEBUG, logging.DEBUG,
f"Detected model class '{model_class}' from filename '{model_name}'", f"Detected model class '{model_class}' from filename '{model_name}'",
min_level=DEBUG_LEVEL.minimal, min_level=DEBUG_LEVEL.minimal,
) )
for m in self.known_models: for m in self.known_models:
env_info["is_" + m] = False env_info.is_flags[m] = False
model_obj = self.models_config.get(m) model_obj = self.models_config.get(m)
model_detect = (model_obj.detect if model_obj else None) or {} model_detect = (model_obj.detect if model_obj else None) or {}
model_detect_for_app: ModelDetectConfig | None = model_detect.get(app) model_detect_for_app: ModelDetectConfig | None = model_detect.get(app.value)
if model_detect_for_app is not None: if model_detect_for_app is not None:
cls_list = model_detect_for_app.class_ or [] cls_list = model_detect_for_app.class_ or []
if model_class and model_class in cls_list: if model_class and model_class in cls_list:
env_info["is_" + m] = True env_info.is_flags[m] = True
elif model_detect_for_app.property is not None and prop_base is not None: elif model_detect_for_app.property is not None and prop_base is not None:
prop = model_detect_for_app.property prop = model_detect_for_app.property
attr = getattr(prop_base, prop, None) attr = getattr(prop_base, prop, None)
if isinstance(attr, bool) and attr: if isinstance(attr, bool) and attr:
env_info["is_" + m] = True env_info.is_flags[m] = True
def __on_model_info_update(self) -> None: def __on_model_info_update(self) -> None:
"""Called when _modelfullname or _modelclass are set via a prompt command.""" """Called when _modelfullname or _modelclass are set via a prompt command."""
@@ -438,7 +459,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
def update( def update(
self, self,
env_info: dict[str, Any], env_info: PPPEnvInfo,
options: PPPStateOptions, options: PPPStateOptions,
wildcards_obj: PPPWildcards, wildcards_obj: PPPWildcards,
extranetwork_mappings_obj: PPPExtraNetworkMappings, extranetwork_mappings_obj: PPPExtraNetworkMappings,
@@ -618,14 +639,21 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
) )
@property @property
def envinfo_hash(self) -> str: def envinfo_hash(self) -> int:
""" """Returns a hash of the environment information for cache-invalidation purposes."""
Generates a hash string based on the environment information. ei = self.state.env_info
ppp_config_key = ei.ppp_config if isinstance(ei.ppp_config, (str, type(None))) else id(ei.ppp_config)
Returns: return hash(
str: A hash string representing the environment information. (
""" ei.app,
return hash(tuple(sorted(self.state.env_info.items()))) ppp_config_key,
ei.model_class,
ei.model_filename,
ei.models_path,
id(ei.property_base),
tuple(sorted(ei.is_flags.items())),
)
)
@property @property
def options_hash(self) -> str: def options_hash(self) -> str:
@@ -661,19 +689,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
vs.set_system(var_name, str(opt_value).split(".", 1)[-1]) vs.set_system(var_name, str(opt_value).split(".", 1)[-1])
# Model related variables # Model related variables
sdchecks = {x: self.state.env_info.get("is_" + x, False) for x in self.known_models} sdchecks = {x: self.state.env_info.is_flags.get(x, False) for x in self.known_models}
# Adding "" as a sentinel that is always True lets next() return "" when no model matches, # Adding "" as a sentinel that is always True lets next() return "" when no model matches,
# giving a well-defined empty-string fallback without a separate None check. # giving a well-defined empty-string fallback without a separate None check.
sdchecks.update({"": True}) sdchecks.update({"": True})
model_name_val = next((k for k, v in sdchecks.items() if v), "") model_name_val = next((k for k, v in sdchecks.items() if v), "")
vs.set_system("_model", model_name_val) vs.set_system("_model", model_name_val)
vs.set_system("_sd", model_name_val) # deprecated vs.set_system("_sd", model_name_val) # deprecated
model_filename = self.state.env_info.get("model_filename", "") model_filename = self.state.env_info.model_filename
vs.set_system("_sdfullname", model_filename) # deprecated vs.set_system("_sdfullname", model_filename) # deprecated
vs.set_system("_modelfullname", model_filename) vs.set_system("_modelfullname", model_filename)
vs.set_system("_sdname", Path(model_filename).name) # deprecated vs.set_system("_sdname", Path(model_filename).name) # deprecated
vs.set_system("_modelname", Path(model_filename).name) vs.set_system("_modelname", Path(model_filename).name)
vs.set_system("_modelclass", self.state.env_info.get("model_class", "")) vs.set_system("_modelclass", self.state.env_info.model_class)
is_models = {} is_models = {}
for model_name, model_type_and_substrings in self.variants_definitions.items(): for model_name, model_type_and_substrings in self.variants_definitions.items():
# A variant is only active when its parent model type is currently loaded # A variant is only active when its parent model type is currently loaded
@@ -701,8 +729,14 @@ 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_pure_" + x, sdchecks[x] and not any(is_models.values()))
vs.set_system("_is_variant_" + x, sdchecks[x] and any(is_models.values())) vs.set_system("_is_variant_" + x, sdchecks[x] and any(is_models.values()))
# special cases # 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_ssd = self.state.env_info.get("is_ssd", False) "_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.is_flags.get("ssd", False)
vs.set_system("_is_ssd", is_ssd) vs.set_system("_is_ssd", is_ssd)
vs.set_system("_is_sdxl_no_ssd", sdchecks.get("sdxl", False) and not is_ssd) vs.set_system("_is_sdxl_no_ssd", sdchecks.get("sdxl", False) and not is_ssd)
# backcompatibility (but the modern one to use would be _is_pure_sdxl) # backcompatibility (but the modern one to use would be _is_pure_sdxl)
@@ -761,7 +795,11 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
self.log(logging.DEBUG, f"BREAK construct {break_replacements[break_processing][0]}") self.log(logging.DEBUG, f"BREAK construct {break_replacements[break_processing][0]}")
elif break_processing == "error": elif break_processing == "error":
if re.search(r"\bBREAK\b", text): if re.search(r"\bBREAK\b", text):
warn_or_stop(self.state, where == -1, "BREAK constructs are not allowed!") warn_or_stop(
self.state,
WARN_STOP_WHERE.negative if where == -1 else WARN_STOP_WHERE.positive,
"BREAK constructs are not allowed!",
)
if self.state.options.cup_ands: if self.state.options.cup_ands:
# collapse ANDs with space after # collapse ANDs with space after
@@ -920,7 +958,11 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
var_keys = sorted(variables_snapshot.keys()) var_keys = sorted(variables_snapshot.keys())
for k in var_keys: for k in var_keys:
entry = variables_snapshot[k] entry = variables_snapshot[k]
if entry.last_echoed_evaluated_value is None and entry.value is not None: if (
entry.last_echoed_evaluated_value is None
and entry.value is not None
and not k.startswith("_output_")
):
unechoed_variables.append(k) unechoed_variables.append(k)
ev = entry.last_echoed_evaluated_value if entry.last_echoed_evaluated_value is not None else entry.value ev = entry.last_echoed_evaluated_value if entry.last_echoed_evaluated_value is not None else entry.value
if ev is not None: if ev is not None:
@@ -941,7 +983,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
# Check for special character sequences that should not be in the result # Check for special character sequences that should not be in the result
compound_prompt = prompt + "\n" + negative_prompt compound_prompt = prompt + "\n" + negative_prompt
found_sequences = re.findall(r"::|\$\$|\$\{|[{}]", compound_prompt) found_sequences = re.findall(r"::|\$\$|\$\{|[{}]|__", compound_prompt)
if found_sequences: if found_sequences:
s = ", ".join(map(lambda x: '"' + x + '"', set(found_sequences))) s = ", ".join(map(lambda x: '"' + x + '"', set(found_sequences)))
warnings.append(f"Probably invalid character sequences: {s}.") warnings.append(f"Probably invalid character sequences: {s}.")
@@ -1014,8 +1056,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
self, self,
prompt: str, prompt: str,
negative_prompt: str, negative_prompt: str,
seed: int, starting_seed: int | list[int],
jobinfo: Any = None, jobinfo: Any = None,
input_vars: dict[str, Any] | None = None,
) -> list[tuple[str, str, dict[str, Any]]]: ) -> list[tuple[str, str, dict[str, Any]]]:
""" """
Process the prompt and negative prompt. Process the prompt and negative prompt.
@@ -1023,7 +1066,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
Args: Args:
prompt (str): The prompt. prompt (str): The prompt.
negative_prompt (str): The negative prompt. negative_prompt (str): The negative prompt.
seed (int): The seed for the random number generator. starting_seed (int | list[int]): The starting 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: Returns:
list[tuple[str, str, dict[str, Any]]]: A list of tuples, each containing the processed prompt, negative prompt, and all variables. list[tuple[str, str, dict[str, Any]]]: A list of tuples, each containing the processed prompt, negative prompt, and all variables.
@@ -1031,15 +1076,19 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
self.state.variables.clear_user() self.state.variables.clear_user()
# We update the input state # We update the input state
# Truncate the seed to the host's configured bit width (-1 because we only want self.state.inputs.seed = (
# positive numbers) so the value stays within the range the host expects clamp_host_bits(self.state.host_config.seed_bits, starting_seed)
# (e.g., 32-bit for SD-WebUI, 64-bit for ComfyUI). if not isinstance(starting_seed, list)
self.state.inputs.seed = int(seed & ((1 << (self.state.host_config.seed_bits - 1)) - 1)) else [clamp_host_bits(self.state.host_config.seed_bits, s) for s in starting_seed]
)
self.state.inputs.pos_prompt = prompt self.state.inputs.pos_prompt = prompt
self.state.inputs.neg_prompt = negative_prompt self.state.inputs.neg_prompt = negative_prompt
self.state.inputs.jobinfo = jobinfo self.state.inputs.jobinfo = jobinfo
# Input related system variables # 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(): for input_name in self.state.inputs.__dict__.keys():
input_value = getattr(self.state.inputs, input_name) input_value = getattr(self.state.inputs, input_name)
var_name = "_input_" + input_name var_name = "_input_" + input_name
@@ -1061,7 +1110,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
filtered_sysvars_inputs = {k: v for k, v in self.state.variables.all_system.items() if k.startswith("_input_")} filtered_sysvars_inputs = {k: v for k, v in self.state.variables.all_system.items() if k.startswith("_input_")}
self.log(logging.INFO, f"Inputs: {filtered_sysvars_inputs}") self.log(logging.INFO, f"Inputs: {filtered_sysvars_inputs}")
rng = np.random.default_rng(self.state.inputs.seed) rng = np.random.default_rng(
self.state.inputs.seed if not isinstance(self.state.inputs.seed, list) else self.state.inputs.seed[0]
)
# Parse both prompts # Parse both prompts
processor = TreeProcessor(self.state, rng, on_model_info_update=self.__on_model_info_update) processor = TreeProcessor(self.state, rng, on_model_info_update=self.__on_model_info_update)
@@ -1094,21 +1145,34 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
final_results: list[tuple[str, str, dict[str, Any]]] = [] final_results: list[tuple[str, str, dict[str, Any]]] = []
for i, r in enumerate(results): for i, r in enumerate(results):
if self.state.options.do_combinatorial: if self.state.options.run_mode == RUN_MODE.combinatorial:
self.log(logging.INFO, f"Combination {i + 1}:") self.log(logging.INFO, f"Combination {i + 1}:")
elif self.state.options.run_mode == RUN_MODE.multiple:
self.log(logging.INFO, f"Result {i + 1}:")
final_results.append(self.__postprocess_result(r)) final_results.append(self.__postprocess_result(r))
if self.state.options.do_combinatorial: if self.state.options.run_mode == RUN_MODE.combinatorial:
self.log(logging.INFO, f"Total combinations: {len(final_results)}") self.log(logging.INFO, f"Total combinations: {len(final_results)}")
if self.state.options.combinatorial_shuffle: if self.state.options.results_shuffle:
rng.shuffle(final_results) rng.shuffle(final_results)
self.log(logging.INFO, "Combinations shuffled") self.log(logging.INFO, "Results shuffled")
return final_results return final_results
def process_prompts_group_start(self): def process_prompts_group_start(self):
"""Start of a prompt processing group.""" """Start of a prompt processing group."""
filtered_sysvars = {k: v for k, v in self.state.variables.all_system.items() if not k.startswith("_input_")} filtered_sysvars = {
self.log(logging.DEBUG, f"System variables: {filtered_sysvars}") k: v for k, v in self.state.variables.all_system.items() if not k.startswith(("_input_", "_output_"))
self.log(logging.INFO, f"Combinatorial: {self.state.options.do_combinatorial}") }
self.log(logging.DEBUG, f"System variables: {filtered_sysvars}", DEBUG_LEVEL.minimal)
self.log(logging.INFO, f"Run mode: {self.state.options.run_mode.name}")
if self.state.options.run_mode == RUN_MODE.combinatorial:
if self.state.options.results_limit > 0:
self.log(logging.INFO, f"Up to {self.state.options.results_limit} combinations")
else:
self.log(logging.INFO, "No combinations limit")
elif self.state.options.run_mode == RUN_MODE.multiple:
if self.state.options.results_limit < 1:
self.state.options.results_limit = 1
self.log(logging.INFO, f"Returning {self.state.options.results_limit} results")
def _expand_filename(self) -> Path: def _expand_filename(self) -> Path:
"""Expand %...% tokens in a filename template and resolve relative paths against the extension logs folder.""" """Expand %...% tokens in a filename template and resolve relative paths against the extension logs folder."""
@@ -1117,7 +1181,7 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
r"%datetime%": now.strftime(r"%Y-%m-%d_%H-%M-%S"), r"%datetime%": now.strftime(r"%Y-%m-%d_%H-%M-%S"),
r"%date%": now.strftime(r"%Y-%m-%d"), r"%date%": now.strftime(r"%Y-%m-%d"),
r"%time%": now.strftime(r"%H-%M-%S"), r"%time%": now.strftime(r"%H-%M-%S"),
r"%host%": str(self.state.env_info.get("app", "")), r"%host%": self.state.env_info.app.value,
} }
result = str(self.state.options.results_file) result = str(self.state.options.results_file)
for token, value in substitutions.items(): for token, value in substitutions.items():
@@ -1145,11 +1209,16 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
"system_variables": { "system_variables": {
k: v k: v
for k, v in all_variables.items() for k, v in all_variables.items()
if k.startswith("_") and not k.startswith("_opt_") and not k.startswith("_input_") if k.startswith("_")
and not k.startswith("_opt_")
and not k.startswith(("_input_", "_output_"))
}, },
"inputs": { "inputs": {
k.removeprefix("_input_"): v for k, v in all_variables.items() if k.startswith("_input_") k.removeprefix("_input_"): v for k, v in all_variables.items() if k.startswith("_input_")
}, },
"outputs": {
k.removeprefix("_output_"): v for k, v in all_variables.items() if k.startswith("_output_")
},
"prompt_results": {"prompt": result_prompt, "negative_prompt": result_neg_prompt}, "prompt_results": {"prompt": result_prompt, "negative_prompt": result_neg_prompt},
"user_variables": {k: v for k, v in all_variables.items() if not k.startswith("_")}, "user_variables": {k: v for k, v in all_variables.items() if not k.startswith("_")},
} }
@@ -1216,8 +1285,9 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
self, self,
original_prompt: str, original_prompt: str,
original_negative_prompt: str, original_negative_prompt: str,
seed: int = -1, starting_seed: int | list[int] = -1,
jobinfo: Any = None, jobinfo: Any = None,
input_vars: dict[str, Any] | None = None,
) -> list[tuple[str, str, dict[str, Any]]]: ) -> list[tuple[str, str, dict[str, Any]]]:
""" """
Initializes the random number generator and processes the prompt and negative prompt. Initializes the random number generator and processes the prompt and negative prompt.
@@ -1225,23 +1295,29 @@ class PromptPostProcessor: # pylint: disable=too-few-public-methods,too-many-in
Args: Args:
original_prompt (str): The original prompt. original_prompt (str): The original prompt.
original_negative_prompt (str): The original negative prompt. original_negative_prompt (str): The original negative prompt.
seed (int): The seed. starting_seed (int | list[int]): The starting seed or list of starting seeds.
jobinfo (Any): Optional job information, available as `_input_jobinfo`. 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: Returns:
list[tuple[str, str, dict[str, Any]]]: A list of tuples containing the processed prompt, negative prompt and all the prompt variables. list[tuple[str, str, dict[str, Any]]]: A list of tuples containing the processed prompt, negative prompt and all the prompt variables.
""" """
results: list[tuple[str, str, dict[str, Any]]] results: list[tuple[str, str, dict[str, Any]]]
try: try:
if seed == -1: if isinstance(starting_seed, list):
seed = np.random.randint(0, 2 ** (self.state.host_config.seed_bits - 1), dtype=np.int64) starting_seed = [
np.random.randint(0, 1 << (self.state.host_config.seed_bits - 1), dtype=np.int64) if s == -1 else s
for s in starting_seed
]
elif starting_seed == -1:
starting_seed = np.random.randint(0, 1 << (self.state.host_config.seed_bits - 1), dtype=np.int64)
prompt = original_prompt prompt = original_prompt
negative_prompt = original_negative_prompt negative_prompt = original_negative_prompt
t1 = time.monotonic_ns() t1 = time.monotonic_ns()
if self.state.cyclical_state.last_prompt_pair != (original_prompt, original_negative_prompt): if self.state.cyclical_state.last_prompt_pair != (original_prompt, original_negative_prompt):
self.state.cyclical_state.reset() self.state.cyclical_state.reset()
self.state.cyclical_state.last_prompt_pair = (original_prompt, original_negative_prompt) 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, starting_seed, jobinfo, input_vars or {})
t2 = time.monotonic_ns() t2 = time.monotonic_ns()
self.log(logging.INFO, f"Process prompt pair time: {(t2 - t1) / 1_000_000_000:.3f} seconds") 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__()}") # self.log(logging.DEBUG,f"Wildcards memory usage: {self.state.wildcards_obj.__sizeof__()}")
+55 -6
View File
@@ -47,6 +47,24 @@ class ONWARNING_CHOICES(Enum):
stop = "stop" stop = "stop"
class RUN_MODE(Enum):
single = "single"
multiple = "multiple"
combinatorial = "combinatorial"
class DEFAULT_SAMPLER(Enum):
random = "random"
cyclical = "cyclical"
class NEXT_SEED(Enum):
randomize = "randomize"
input = "input"
increment = "increment"
decrement = "decrement"
# ------------------- Host configuration ------------------- # ------------------- Host configuration -------------------
AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"] AttentionOption = Literal["ok", "parentheses", "disable", "remove", "error"]
@@ -209,10 +227,13 @@ class PPPStateOptions:
cup_merge_attention: bool = True cup_merge_attention: bool = True
cup_remove_extranetwork_tags: bool = False cup_remove_extranetwork_tags: bool = False
strict_operators: bool = True strict_operators: bool = True
do_combinatorial: bool = False
combinatorial_shuffle: bool = False
combinatorial_limit: int = 100 # 0 = no limit
results_file: str = "" # empty = disabled; supports %datetime%, %date%, %time%, %host% tokens results_file: str = "" # empty = disabled; supports %datetime%, %date%, %time%, %host% tokens
run_mode: RUN_MODE = RUN_MODE.single
results_limit: int = 100 # 0 = no limit
results_shuffle: bool = False
comb_random_fixed: bool = True # if True, the random sampler will be fixed across all DFS runs
default_sampler: DEFAULT_SAMPLER = DEFAULT_SAMPLER.random
next_seed: NEXT_SEED = NEXT_SEED.randomize # how to determine the next seed for each prompt
def __post_init__(self): def __post_init__(self):
if not self.cup_do_cleanup: if not self.cup_do_cleanup:
@@ -235,7 +256,7 @@ class PPPStateOptions:
class PPPStateInputs: class PPPStateInputs:
"""Structured inputs for a single prompt processing call.""" """Structured inputs for a single prompt processing call."""
seed: int = -1 seed: int | list[int] = -1
pos_prompt: str = "" pos_prompt: str = ""
neg_prompt: str = "" neg_prompt: str = ""
jobinfo: Any = None jobinfo: Any = None
@@ -272,12 +293,30 @@ class CyclicalSamplerState:
self.last_prompt_pair = None self.last_prompt_pair = None
@dataclass
class PPPEnvInfo:
"""Environment and model information passed to PPP at construction time."""
app: SUPPORTED_APPS = SUPPORTED_APPS.a1111
ppp_config: str | dict | None = None
model_class: str = ""
model_filename: str = ""
property_base: Any = None
models_path: str = ""
_is_flags: dict[str, bool] = field(default_factory=dict, init=False, repr=False)
@property
def is_flags(self) -> dict[str, bool]:
"""Boolean model-detection flags keyed by model name (e.g. 'sdxl' -> True)."""
return self._is_flags
@dataclass(frozen=True) @dataclass(frozen=True)
class PPPState: class PPPState:
"""State object passed to various PPP components during prompt processing.""" """State object passed to various PPP components during prompt processing."""
logger: Logger logger: Logger
env_info: dict[str, Any] = field(default_factory=dict) env_info: PPPEnvInfo = field(default_factory=PPPEnvInfo)
host_config: HostConfig = field(default_factory=HostConfig) host_config: HostConfig = field(default_factory=HostConfig)
options: PPPStateOptions = field(default_factory=PPPStateOptions) options: PPPStateOptions = field(default_factory=PPPStateOptions)
inputs: PPPStateInputs = field(default_factory=PPPStateInputs) inputs: PPPStateInputs = field(default_factory=PPPStateInputs)
@@ -288,7 +327,17 @@ class PPPState:
cyclical_state: CyclicalSamplerState = field(default_factory=CyclicalSamplerState) cyclical_state: CyclicalSamplerState = field(default_factory=CyclicalSamplerState)
class PPPInterrupt(Exception): class PPPException(Exception):
"""
Custom exception to handle exceptions in the PromptPostProcessor.
"""
def __init__(self, message: str = "An error occurred during prompt processing."):
super().__init__(message)
self.message = message
class PPPInterrupt(PPPException):
""" """
Custom exception to handle interruptions in the PromptPostProcessor. Custom exception to handle interruptions in the PromptPostProcessor.
This exception can be raised to stop the processing of prompts. This exception can be raised to stop the processing of prompts.
+160 -59
View File
@@ -1,17 +1,28 @@
if __name__ == "__main__": if __name__ == "__main__":
raise SystemExit("This script must be run from ComfyUI") raise SystemExit("This script must be run from ComfyUI")
# pylint: disable=wrong-import-position,wrong-import-order
from datetime import datetime from datetime import datetime
import logging import logging
import os import os
from pathlib import Path from pathlib import Path
from typing import Any from typing import Any
import folder_paths # type: ignore import folder_paths # type: ignore # pylint: disable=import-error
import nodes # type: ignore import nodes # type: ignore # pylint: disable=import-error
from ppp import PromptPostProcessor from ppp import PromptPostProcessor
from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, PPPStateOptions from ppp_classes import (
DEFAULT_SAMPLER,
IFWILDCARDS_CHOICES,
ONWARNING_CHOICES,
PPPEnvInfo,
SUPPORTED_APPS,
PPPException,
RUN_MODE,
NEXT_SEED,
PPPStateOptions,
)
from ppp_common import get_model_class_from_filename, load_grammar from ppp_common import get_model_class_from_filename, load_grammar
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log
from ppp_utils import escape_single_quotes from ppp_utils import escape_single_quotes
@@ -190,29 +201,20 @@ class PromptPostProcessorComfyUINode:
"label_off": "No", "label_off": "No",
}, },
), ),
"do_combinatorial": ( "results_file": (
"BOOLEAN", "STRING",
{ {
"default": PromptPostProcessor.DEFAULT_DO_COMBINATORIAL, "default": PromptPostProcessor.DEFAULT_RESULTS_FILE,
"tooltip": "Enable combinatorial mode", "tooltip": r"Filename to save processing results. Supports %datetime%, %date%, %time%, %host% tokens. Empty = disabled.",
"label_on": "Yes", "dynamicPrompts": False,
"label_off": "No",
}, },
), ),
"combinatorial_shuffle": ( "run_mode": (
"BOOLEAN", "COMBO",
{ {
"default": PromptPostProcessor.DEFAULT_COMBINATORIAL_SHUFFLE, "options": [e.value for e in RUN_MODE],
"tooltip": "Shuffle the combinatorial results", "default": PromptPostProcessor.DEFAULT_RUN_MODE,
"label_on": "Yes", "tooltip": "Run mode",
"label_off": "No",
},
),
"combinatorial_limit": (
"INT",
{
"default": PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT,
"tooltip": "Limit for combinatorial mode",
}, },
), ),
"wc_options": ( "wc_options": (
@@ -243,12 +245,11 @@ class PromptPostProcessorComfyUINode:
"tooltip": "ExtraNetworks mapping options", "tooltip": "ExtraNetworks mapping options",
}, },
), ),
"results_file": ( "rm_options": (
"STRING", "PPP_OPTIONS_RM",
{ {
"default": PromptPostProcessor.DEFAULT_RESULTS_FILE, "default": None,
"tooltip": r"Filename to save processing results. Supports %datetime%, %date%, %time%, %host% tokens. Empty = disabled.", "tooltip": "Run mode options",
"dynamicPrompts": False,
}, },
), ),
}, },
@@ -283,9 +284,9 @@ class PromptPostProcessorComfyUINode:
"variables", "variables",
) )
OUTPUT_TOOLTIPS = ( OUTPUT_TOOLTIPS = (
"Processed positive prompt (list of prompts if combinatorial mode is enabled)", "Processed positive prompt (list of prompts if combinatorial/multiple mode is enabled)",
"Processed negative prompt (list of prompts if combinatorial mode is enabled)", "Processed negative prompt (list of prompts if combinatorial/multiple mode is enabled)",
"Output variables (list of dictionaries if combinatorial mode is enabled)", "Output variables (list of dictionaries if combinatorial/multiple mode is enabled)",
) )
FUNCTION = "process" FUNCTION = "process"
@@ -304,27 +305,34 @@ class PromptPostProcessorComfyUINode:
seed, seed,
debug_level, debug_level,
on_warnings, on_warnings,
strict_operators,
process_wildcards, process_wildcards,
do_cleanup, do_cleanup,
cleanup_variables, cleanup_variables,
do_combinatorial, results_file,
combinatorial_shuffle, run_mode,
combinatorial_limit,
model=None, model=None,
wc_options=None, wc_options=None,
stn_options=None, stn_options=None,
cup_options=None, cup_options=None,
en_options=None, en_options=None,
strict_operators=None, rm_options=None,
results_file=None,
): ):
modelclass = ( modelclass = (
model.model.model_config.__class__.__name__ if model is not None and not isinstance(model, str) else model model.model.model_config.__class__.__name__ if model is not None and not isinstance(model, str) else model
) or "" ) or ""
if modelname == "(none)": if modelname == "(none)":
modelname = "" modelname = ""
if modelclass == "": if modelclass == "" and modelname != "":
modelclass = get_model_class_from_filename(modelname) try:
modelclass = get_model_class_from_filename(Path(modelname))
except PPPException as e:
log(
self.logger,
DEBUG_LEVEL.minimal,
logging.WARNING,
f"Could not detect model class from filename '{modelname}': {e}",
)
if modelclass: if modelclass:
log( log(
self.logger, self.logger,
@@ -337,7 +345,7 @@ class PromptPostProcessorComfyUINode:
self.logger, self.logger,
DEBUG_LEVEL.minimal, DEBUG_LEVEL.minimal,
logging.WARNING, logging.WARNING,
"Model class was not provided. System model variables will not be properly set.", "Model class was not provided nor detected. System model variables will not be properly set.",
) )
if modelname == "": if modelname == "":
log( log(
@@ -347,24 +355,26 @@ class PromptPostProcessorComfyUINode:
"Modelname was not provided. System model and variant variables will not be properly set.", "Modelname was not provided. System model and variant variables will not be properly set.",
) )
# model class values in ComfyUI\comfy\supported_models.py # model class values in ComfyUI\comfy\supported_models.py
env_info = { env_info = PPPEnvInfo(
"app": SUPPORTED_APPS.comfyui.value, app=SUPPORTED_APPS.comfyui,
"models_path": folder_paths.models_dir, models_path=folder_paths.models_dir,
"model_filename": modelname or "", # path is relative to checkpoints folder model_filename=modelname or "", # path is relative to checkpoints folder
"model_class": modelclass, model_class=modelclass,
"property_base": None, property_base=None,
} )
wildcards_folders = _resolve_wildcards_folders(wc_options["wc_wildcards_folders"] if wc_options else "") wildcards_folders = _resolve_wildcards_folders(wc_options["wc_wildcards_folders"] if wc_options else "")
enmappings_folders = _resolve_enmappings_folders(en_options["en_mappings_folders"] if en_options else "") enmappings_folders = _resolve_enmappings_folders(en_options["en_mappings_folders"] if en_options else "")
options = PPPStateOptions( options = PPPStateOptions(
debug_level=DEBUG_LEVEL(debug_level), debug_level=DEBUG_LEVEL(debug_level),
on_warning=ONWARNING_CHOICES(on_warnings) if on_warnings else PromptPostProcessor.DEFAULT_ON_WARNING, on_warning=ONWARNING_CHOICES(on_warnings if on_warnings else PromptPostProcessor.DEFAULT_ON_WARNING),
strict_operators=( strict_operators=(
strict_operators if strict_operators is not None else PromptPostProcessor.DEFAULT_STRICT_OPERATORS strict_operators if strict_operators is not None else PromptPostProcessor.DEFAULT_STRICT_OPERATORS
), ),
process_wildcards=process_wildcards, process_wildcards=process_wildcards,
if_wildcards=(wc_options["wc_if_wildcards"] if wc_options else IFWILDCARDS_CHOICES.stop.value), if_wildcards=IFWILDCARDS_CHOICES(
wc_options["wc_if_wildcards"] if wc_options else PromptPostProcessor.DEFAULT_IF_WILDCARDS
),
choice_separator=( choice_separator=(
wc_options["wc_choice_separator"] if wc_options else PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR wc_options["wc_choice_separator"] if wc_options else PromptPostProcessor.DEFAULT_CHOICE_SEPARATOR
), ),
@@ -415,10 +425,19 @@ class PromptPostProcessorComfyUINode:
if cup_options if cup_options
else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS else PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS
), ),
do_combinatorial=do_combinatorial, run_mode=RUN_MODE(run_mode if run_mode else PromptPostProcessor.DEFAULT_RUN_MODE),
combinatorial_shuffle=combinatorial_shuffle, results_file=results_file,
combinatorial_limit=combinatorial_limit, results_shuffle=(
results_file=results_file or "", rm_options["results_shuffle"] if rm_options else PromptPostProcessor.DEFAULT_RESULTS_SHUFFLE
),
results_limit=rm_options["results_limit"] if rm_options else PromptPostProcessor.DEFAULT_RESULTS_LIMIT,
comb_random_fixed=(
rm_options["comb_random_fixed"] if rm_options else PromptPostProcessor.DEFAULT_COMB_RANDOM_FIXED
),
default_sampler=DEFAULT_SAMPLER(
rm_options["default_sampler"] if rm_options else PromptPostProcessor.DEFAULT_DEFAULT_SAMPLER
),
next_seed=NEXT_SEED(rm_options["next_seed"] if rm_options else PromptPostProcessor.DEFAULT_NEXT_SEED),
) )
self.wildcards_obj.refresh_wildcards( self.wildcards_obj.refresh_wildcards(
options.debug_level, options.debug_level,
@@ -455,13 +474,94 @@ class PromptPostProcessorComfyUINode:
jobinfo={"job_timestamp": datetime.now().isoformat()}, jobinfo={"job_timestamp": datetime.now().isoformat()},
) )
self.ppp.process_prompts_group_end() self.ppp.process_prompts_group_end()
return tuple(zip(*results)) # unzip the list of tuples into tuple of lists return tuple(zip(*results)) # unzip the list of tuples into tuple of lists
def interrupt(self): def interrupt(self):
nodes.interrupt_processing(True) nodes.interrupt_processing(True)
class PromptPostProcessorRunModeOptionsComfyUINode:
"""
Node for run mode options.
"""
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls):
return {
"optional": {
"results_limit": (
"INT",
{
"default": PromptPostProcessor.DEFAULT_RESULTS_LIMIT,
"tooltip": "Limit for combinatorial/multiple mode",
"min": 0,
},
),
"results_shuffle": (
"BOOLEAN",
{
"default": PromptPostProcessor.DEFAULT_RESULTS_SHUFFLE,
"tooltip": "Shuffle the results",
"label_on": "Yes",
"label_off": "No",
},
),
"comb_random_fixed": (
"BOOLEAN",
{
"default": PromptPostProcessor.DEFAULT_COMB_RANDOM_FIXED,
"tooltip": "Fix the value of any specified random samplers across all combinations in combinatorial mode",
"label_on": "Yes",
"label_off": "No",
},
),
"default_sampler": (
"COMBO",
{
"options": [ds.value for ds in DEFAULT_SAMPLER],
"default": PromptPostProcessor.DEFAULT_DEFAULT_SAMPLER,
"tooltip": "Default choice sampler",
},
),
"next_seed": (
"COMBO",
{
"options": [e.value for e in NEXT_SEED],
"default": PromptPostProcessor.DEFAULT_NEXT_SEED,
"tooltip": "Next seed strategy",
},
),
},
}
RETURN_TYPES = ("PPP_OPTIONS_RM",)
RETURN_NAMES = ("options",)
FUNCTION = "process"
CATEGORY = "ACB"
def process(
self,
results_limit: int,
results_shuffle: bool,
comb_random_fixed: bool,
default_sampler: str,
next_seed: str,
):
options = {
"results_limit": results_limit,
"results_shuffle": results_shuffle,
"comb_random_fixed": comb_random_fixed,
"default_sampler": default_sampler,
"next_seed": next_seed,
}
return (options,)
class PromptPostProcessorWildcardOptionsComfyUINode: class PromptPostProcessorWildcardOptionsComfyUINode:
""" """
Node for wildcard options. Node for wildcard options.
@@ -907,13 +1007,13 @@ class PromptPostProcessorWildcardConcatComfyUINode:
lf = PromptPostProcessorLogFactory() lf = PromptPostProcessorLogFactory()
cls._ppp = PromptPostProcessor( cls._ppp = PromptPostProcessor(
lf.log, lf.log,
{ PPPEnvInfo(
"app": SUPPORTED_APPS.comfyui.value, app=SUPPORTED_APPS.comfyui,
"models_path": folder_paths.models_dir, models_path=folder_paths.models_dir,
"model_filename": "", model_filename="",
"model_class": "", model_class="",
"property_base": None, property_base=None,
}, ),
PPPStateOptions(debug_level=DEBUG_LEVEL.minimal), PPPStateOptions(debug_level=DEBUG_LEVEL.minimal),
wildcards_obj=PPPWildcards(lf.log), wildcards_obj=PPPWildcards(lf.log),
) )
@@ -1031,6 +1131,7 @@ class PromptPostProcessorWildcardConcatComfyUINode:
try: try:
# pylint: disable=import-error
from server import PromptServer # type: ignore from server import PromptServer # type: ignore
from aiohttp import web as _aiohttp_web # type: ignore from aiohttp import web as _aiohttp_web # type: ignore
+85 -20
View File
@@ -1,5 +1,6 @@
import ast import ast
import csv import csv
from enum import Enum
from functools import reduce from functools import reduce
import logging import logging
from pathlib import Path from pathlib import Path
@@ -10,7 +11,7 @@ import lark
from ruamel.yaml import YAML as _YAML from ruamel.yaml import YAML as _YAML
from ppp_logging import log from ppp_logging import log
from ppp_classes import ONWARNING_CHOICES, PPPInterrupt, PPPState from ppp_classes import ONWARNING_CHOICES, PPPException, PPPInterrupt, PPPState
from ppp_utils import escape_single_quotes, format_output from ppp_utils import escape_single_quotes, format_output
@@ -82,13 +83,19 @@ def parse_prompt(
return parsed_prompt return parsed_prompt
def warn_or_stop(state: PPPState, is_negative: bool, message: str, e: Exception = None): class WARN_STOP_WHERE(Enum):
none = 0
positive = 1
negative = 2
def warn_or_stop(state: PPPState, where: WARN_STOP_WHERE, message: str, e: Exception = None):
INVALID_CONTENT_STOP = "INVALID CONTENT! {0}\nBREAK " INVALID_CONTENT_STOP = "INVALID CONTENT! {0}\nBREAK "
if state.options.on_warning == ONWARNING_CHOICES.stop: if state.options.on_warning == ONWARNING_CHOICES.stop:
raise PPPInterrupt( raise PPPInterrupt(
message, message,
INVALID_CONTENT_STOP.format(message) if not is_negative else "", INVALID_CONTENT_STOP.format(message) if where == WARN_STOP_WHERE.positive else "",
INVALID_CONTENT_STOP.format(message) if is_negative else "", INVALID_CONTENT_STOP.format(message) if where == WARN_STOP_WHERE.negative else "",
) from e ) from e
log(state.logger, state.options.debug_level, logging.WARNING, format_output(message)) log(state.logger, state.options.debug_level, logging.WARNING, format_output(message))
@@ -203,28 +210,43 @@ def preprocess_grammar(grammar_content: str, options: dict[str, bool], logger: l
return "\n".join(result_lines) return "\n".join(result_lines)
def get_model_class_from_filename(filename: str) -> str: def get_model_config_from_filename(filename: Path) -> object | None:
"""
Attempts to detect the model class from the given filename by inspecting the file header.
Currently only supports ComfyUI models in .safetensors format.
The path must be relative to a model folder.
"""
# pylint: disable=import-outside-toplevel
try: try:
import folder_paths # type: ignore import folder_paths # type: ignore
import comfy.utils # type: ignore import comfy.utils # type: ignore
import comfy.model_detection as model_detection # type: ignore import comfy.model_detection as model_detection # type: ignore
except ImportError: except ImportError as e:
return "" raise PPPException(f"Error detecting class from '{filename}': {e}") from e
import json import json
if not filename: if not filename:
return "" return None
full_path = (
folder_paths.get_full_path("diffusion_models", filename) path_keys = ["diffusion_models", "checkpoints", "unet"]
or folder_paths.get_full_path("checkpoints", filename) full_path: Path | None = None
or folder_paths.get_full_path("unet", filename) if filename.is_absolute():
) full_path = filename
if not full_path or not full_path.lower().endswith((".safetensors", ".sft")): else:
return "" base_folders = []
for key in path_keys:
base_folders.extend(folder_paths.get_folder_paths(key))
for base in base_folders:
fname: Path = base / filename
if fname.exists():
full_path = fname
break
if not full_path or not full_path.suffix.lower() in (".safetensors", ".sft"):
return None
try: try:
header_bytes = comfy.utils.safetensors_header(full_path) header_bytes = comfy.utils.safetensors_header(full_path)
if header_bytes is None: if header_bytes is None:
return "" raise PPPException(f"Error detecting class from '{full_path}': no header")
header = json.loads(header_bytes) header = json.loads(header_bytes)
# model_config_from_unet only inspects tensor shapes, not actual data. # model_config_from_unet only inspects tensor shapes, not actual data.
@@ -240,10 +262,39 @@ def get_model_class_from_filename(filename: str) -> str:
mock_sd = {k: _ShapeProxy(v["shape"]) for k, v in header.items() if k != "__metadata__" and "shape" in v} mock_sd = {k: _ShapeProxy(v["shape"]) for k, v in header.items() if k != "__metadata__" and "shape" in v}
prefix = model_detection.unet_prefix_from_state_dict(mock_sd) prefix = model_detection.unet_prefix_from_state_dict(mock_sd)
config = model_detection.model_config_from_unet(mock_sd, prefix) config = model_detection.model_config_from_unet(mock_sd, prefix, True)
return config.__class__.__name__ if config else "" if not config:
except Exception: # pylint: disable=broad-except mock_sd, metadata = comfy.utils.convert_old_quants(mock_sd, "", metadata=None)
# Allow loading unets from checkpoint files
diffusion_model_prefix = model_detection.unet_prefix_from_state_dict(mock_sd)
temp_sd = comfy.utils.state_dict_prefix_replace(mock_sd, {diffusion_model_prefix: ""}, filter_keys=True)
if len(temp_sd) > 0:
mock_sd, metadata = comfy.utils.convert_old_quants(temp_sd, "", metadata=metadata)
config = model_detection.model_config_from_unet(mock_sd, "", metadata=metadata)
if config is None:
mock_sd = model_detection.convert_diffusers_mmdit(mock_sd, "")
if mock_sd is not None: # diffusers mmdit
config = model_detection.model_config_from_unet(mock_sd, "")
else: # diffusers unet
config = model_detection.model_config_from_diffusers_unet(mock_sd)
if not config:
raise PPPException(f"Error detecting class from file '{full_path}': no config found")
return config
except PPPException:
raise
except Exception as e: # pylint: disable=broad-except
raise PPPException(f"Error detecting class from file '{full_path}': {e}") from e
def get_model_class_from_filename(filename: Path) -> str:
config = get_model_config_from_filename(filename)
if not config:
return "" return ""
c = config.__class__.__name__
if not c:
raise PPPException(f"Error detecting class from '{filename}': config has no class name {config}")
return c
def sanitize_wc_name(name: str) -> str: def sanitize_wc_name(name: str) -> str:
@@ -288,7 +339,7 @@ def convert_sdnext_styles_to_wildcard(inp: Path, out: Path):
files = inp.glob("*.json") if inp.is_dir() else [inp] files = inp.glob("*.json") if inp.is_dir() else [inp]
for file in files: for file in files:
with open(file, "r", encoding="utf-8-sig") as f: with open(file, "r", encoding="utf-8-sig") as f:
data = _YAML(typ='safe').load(f) data = _YAML(typ="safe").load(f)
if not isinstance(data, list): if not isinstance(data, list):
continue continue
wildcards[file] = {} wildcards[file] = {}
@@ -309,3 +360,17 @@ def convert_sdnext_styles_to_wildcard(inp: Path, out: Path):
f.write(f"# Converted from {name}\n") f.write(f"# Converted from {name}\n")
yaml_writer = _YAML() yaml_writer = _YAML()
yaml_writer.dump(wcs, f) yaml_writer.dump(wcs, f)
def clamp_host_bits(bits: int, seed: int) -> int:
"""
Clamp the seed to the host's configured bit width so the value stays within the range the host expects (and positive).
Args:
bits (int): The host's configured bit width.
seed (int): The seed to clamp.
Returns:
int: The clamped seed.
"""
return int(seed & ((1 << bits ) - 1))
+64 -8
View File
@@ -82,7 +82,7 @@ hosts:
alternation: error alternation: error
and: comma and: comma
break: comma break: comma
seed_bits: 64 seed_bits: 53 # and not 64, because of frontend JavaScript limits
# Supported base models, variants, and options # Supported base models, variants, and options
# Check supported models for each host in: # Check supported models for each host in:
@@ -101,7 +101,7 @@ models:
a1111: { property: "is_sd1" } a1111: { property: "is_sd1" }
forge: { property: "is_sd1", class: ["SD15", "SD15_instructpix2pix"] } forge: { property: "is_sd1", class: ["SD15", "SD15_instructpix2pix"] }
forgeneo: { property: "is_sd1", class: ["SD15"] } forgeneo: { property: "is_sd1", class: ["SD15"] }
reforge: { property: "is_sd1" } reforge: { property: "is_sd1", class: ["SD15", "SD15_instructpix2pix"] }
sdnext: { class: ["LatentDiffusion", "StableDiffusionPipeline", "StableDiffusionInpaintPipeline", "StableDiffusionInstructPix2PixPipeline", "StableDiffusionUpscalePipeline"] } # LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend; cannot differentiate SD1 and SD2, we set both to True sdnext: { class: ["LatentDiffusion", "StableDiffusionPipeline", "StableDiffusionInpaintPipeline", "StableDiffusionInstructPix2PixPipeline", "StableDiffusionUpscalePipeline"] } # LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend; cannot differentiate SD1 and SD2, we set both to True
comfyui: { class: ["SD15", "SD15_instructpix2pix"] } comfyui: { class: ["SD15", "SD15_instructpix2pix"] }
sd2: # Stable Diffusion 2 sd2: # Stable Diffusion 2
@@ -109,7 +109,7 @@ models:
a1111: { property: "is_sd2" } a1111: { property: "is_sd2" }
forge: { property: "is_sd2", class: ["SD20", "SD21UnclipL", "SD21UnclipH"] } forge: { property: "is_sd2", class: ["SD20", "SD21UnclipL", "SD21UnclipH"] }
forgeneo: null forgeneo: null
reforge: { property: "is_sd2" } reforge: { property: "is_sd2", class: ["SD20", "SD21UnclipL", "SD21UnclipH"] }
sdnext: { class: ["LatentDiffusion", "StableDiffusionPipeline", "StableDiffusionInpaintPipeline", "StableDiffusionInstructPix2PixPipeline", "StableDiffusionUpscalePipeline"] } # cannot differentiate SD1 and SD2, we set both to True; LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend sdnext: { class: ["LatentDiffusion", "StableDiffusionPipeline", "StableDiffusionInpaintPipeline", "StableDiffusionInstructPix2PixPipeline", "StableDiffusionUpscalePipeline"] } # cannot differentiate SD1 and SD2, we set both to True; LatentDiffusion is for the original backend, StableDiffusionPipeline is for the diffusers backend
comfyui: { class: ["SD20", "SD21UnclipL", "SD21UnclipH", "LotusD"] } comfyui: { class: ["SD20", "SD21UnclipL", "SD21UnclipH", "LotusD"] }
ssd: # Segmind Stable Diffusion 1B ssd: # Segmind Stable Diffusion 1B
@@ -117,7 +117,7 @@ models:
a1111: { property: "is_ssd" } a1111: { property: "is_ssd" }
forge: { class: ["SSD1B"] } forge: { class: ["SSD1B"] }
forgeneo: null forgeneo: null
reforge: { property: "is_ssd" } reforge: { property: "is_ssd", class: ["SSD1B"] }
sdnext: null sdnext: null
comfyui: { class: ["SSD1B"]} comfyui: { class: ["SSD1B"]}
sdxl: # Stable Diffusion XL sdxl: # Stable Diffusion XL
@@ -125,7 +125,7 @@ models:
a1111: { property: "is_sdxl" } a1111: { property: "is_sdxl" }
forge: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] } forge: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
forgeneo: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner"] } forgeneo: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner"] }
reforge: { property: "is_sdxl" } reforge: { property: "is_sdxl", class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
sdnext: { class: ["StableDiffusionXLPipeline", "StableDiffusionXLImg2ImgPipeline", "StableDiffusionXLInpaintPipeline", "StableDiffusionXLInstructPix2PixPipeline"] } sdnext: { class: ["StableDiffusionXLPipeline", "StableDiffusionXLImg2ImgPipeline", "StableDiffusionXLInpaintPipeline", "StableDiffusionXLInstructPix2PixPipeline"] }
comfyui: { class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] } comfyui: { class: ["SDXL", "SDXLRefiner", "SDXL_instructpix2pix", "Segmind_Vega", "KOALA_700M", "KOALA_1B"] }
variants: variants:
@@ -142,7 +142,7 @@ models:
a1111: { property: "is_sd3" } a1111: { property: "is_sd3" }
forge: { property: "is_sd3", class: ["SD3"] } forge: { property: "is_sd3", class: ["SD3"] }
forgeneo: null forgeneo: null
reforge: { property: "is_sd3" } reforge: { property: "is_sd3", class: ["SD3"] }
sdnext: { class: ["StableDiffusion3Pipeline"] } sdnext: { class: ["StableDiffusion3Pipeline"] }
comfyui: { class: ["SD3"] } comfyui: { class: ["SD3"] }
flux: # Flux 1 flux: # Flux 1
@@ -190,7 +190,7 @@ models:
a1111: null a1111: null
forge: null forge: null
forgeneo: null forgeneo: null
reforge: null reforge: { class: ["LTXV"] }
sdnext: null sdnext: null
comfyui: { class: ["LTXV", "LTXAV"] } comfyui: { class: ["LTXV", "LTXAV"] }
cosmos: # Cosmos cosmos: # Cosmos
@@ -248,7 +248,7 @@ models:
forgeneo: { class: ["WAN21_T2V", "WAN21_I2V"] } forgeneo: { class: ["WAN21_T2V", "WAN21_I2V"] }
reforge: { class: ["WAN21_T2V", "WAN21_I2V", "WAN21_FunControl2V", "WAN21_Camera", "WAN22_Camera", "WAN21_Vace", "WAN21_HuMo", "WAN22_S2V", "WAN22_Animate", "WAN22_T2V"] } reforge: { class: ["WAN21_T2V", "WAN21_I2V", "WAN21_FunControl2V", "WAN21_Camera", "WAN22_Camera", "WAN21_Vace", "WAN21_HuMo", "WAN22_S2V", "WAN22_Animate", "WAN22_T2V"] }
sdnext: { class: ["WanPipeline"] } sdnext: { class: ["WanPipeline"] }
comfyui: { class: ["WAN21_T2V", "WAN21_I2V", "WAN21_FunControl2V", "WAN21_Camera", "WAN22_Camera", "WAN21_Vace", "WAN21_HuMo", "WAN22_S2V", "WAN22_Animate", "WAN22_T2V", "WAN21_FlowRVS", "WAN21_SCAIL", "WAN22_WanDancer"] } comfyui: { class: ["WAN21_T2V", "WAN21_I2V", "WAN21_FunControl2V", "WAN21_Camera", "WAN22_Camera", "WAN21_Vace", "WAN21_HuMo", "WAN22_S2V", "WAN22_Animate", "WAN22_T2V", "WAN21_FlowRVS", "WAN21_SCAIL", "WAN22_WanDancer", "WAN21_CausalAR_T2V", "WAN21_SCAIL2", "WAN_Animate2"] }
hidream: # HiDream hidream: # HiDream
detect: detect:
a1111: null a1111: null
@@ -337,3 +337,59 @@ models:
reforge: null reforge: null
sdnext: null sdnext: null
comfyui: { class: ["CogVideoX_T2V", "CogVideoX_I2V", "CogVideoX_Inpaint"] } comfyui: { class: ["CogVideoX_T2V", "CogVideoX_I2V", "CogVideoX_Inpaint"] }
krea2: # Krea2
detect:
a1111: null
forge: null
forgeneo: { class: ["Krea2"] }
reforge: null
sdnext: null
comfyui: { class: ["Krea2"] }
ideogram4: # Ideogram4
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["Ideogram4"] }
lens: # Lens
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["Lens"] }
boogu: # Boogu
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["Boogu"] }
mageflow: # MageFlow
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["MageFlow"] }
pid: # PiD
detect:
a1111: null
forge: null
forgeneo: { class: ["PiD"] }
reforge: null
sdnext: null
comfyui: { class: ["PiD"] }
minimaxh3: # MiniMaxH3
detect:
a1111: null
forge: null
forgeneo: null
reforge: null
sdnext: null
comfyui: { class: ["MiniMaxH3"] }
+3 -3
View File
@@ -218,7 +218,7 @@ class PPPExtraNetworkMappings:
enmappings_input = enmappings_input.strip() enmappings_input = enmappings_input.strip()
if enmappings_input != "": if enmappings_input != "":
try: try:
content = _YAML(typ='safe').load(enmappings_input) content = _YAML(typ="safe").load(enmappings_input)
except _YAMLError as e: except _YAMLError as e:
log( log(
self.__logger, self.__logger,
@@ -298,7 +298,7 @@ class PPPExtraNetworkMappings:
try: try:
try: try:
with open(full_path, "r", encoding="utf-8") as file: with open(full_path, "r", encoding="utf-8") as file:
content = _YAML(typ='safe').load(file) content = _YAML(typ="safe").load(file)
except: # pylint: disable=bare-except except: # pylint: disable=bare-except
log( log(
self.__logger, self.__logger,
@@ -307,7 +307,7 @@ class PPPExtraNetworkMappings:
f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...", f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...",
) )
with open(full_path, "r", encoding="windows-1252") as file: with open(full_path, "r", encoding="windows-1252") as file:
content = _YAML(typ='safe').load(file) content = _YAML(typ="safe").load(file)
self.__add_extranetwork_mapping(content, full_path) self.__add_extranetwork_mapping(content, full_path)
except Exception as e: # pylint: disable=broad-except except Exception as e: # pylint: disable=broad-except
log( log(
+222 -109
View File
@@ -11,11 +11,11 @@ from typing import Callable, Optional
import lark import lark
import numpy as np import numpy as np
from ppp_classes import IFWILDCARDS_CHOICES, SUPPORTED_APPS, PPPState from ppp_classes import DEFAULT_SAMPLER, IFWILDCARDS_CHOICES, NEXT_SEED, RUN_MODE, SUPPORTED_APPS, PPPState
from ppp_enmappings import PPPENMappingVariant from ppp_enmappings import PPPENMappingVariant
from ppp_logging import DEBUG_LEVEL, log from ppp_logging import DEBUG_LEVEL, log
from ppp_utils import escape_single_quotes, repr_value from ppp_utils import escape_single_quotes, repr_value
from ppp_common import parse_prompt, warn_or_stop from ppp_common import WARN_STOP_WHERE, clamp_host_bits, parse_prompt, warn_or_stop
from ppp_variables import ScalarValue, VariableEntry from ppp_variables import ScalarValue, VariableEntry
from ppp_wildcards import PPPWildcard from ppp_wildcards import PPPWildcard
@@ -62,6 +62,9 @@ class TreeProcessor(lark.visitors.Interpreter):
self.__on_model_info_update = on_model_info_update self.__on_model_info_update = on_model_info_update
self.__debug_level = state.options.debug_level self.__debug_level = state.options.debug_level
self.__rng = rng self.__rng = rng
self.__current_input_seed = (
self.state.inputs.seed if not isinstance(self.state.inputs.seed, list) else self.state.inputs.seed[0]
)
self.__shell: list[TreeProcessor.AccumulatedShell] = [] # type: ignore self.__shell: list[TreeProcessor.AccumulatedShell] = [] # type: ignore
self.__negtags: list[TreeProcessor.NegTag] = [] # type: ignore self.__negtags: list[TreeProcessor.NegTag] = [] # type: ignore
self.__already_processed: list[str] = [] self.__already_processed: list[str] = []
@@ -71,20 +74,25 @@ class TreeProcessor(lark.visitors.Interpreter):
self.__add_at: dict[str, list] = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []} self.__add_at: dict[str, list] = {"start": [], "insertion_point": [[] for _ in range(10)], "end": []}
self.__insertion_at: list[tuple[int, int]] = [None for _ in range(10)] self.__insertion_at: list[tuple[int, int]] = [None for _ in range(10)]
self.__detectedWildcards: list[tuple[str, bool]] = [] self.__detectedWildcards: list[tuple[str, bool]] = []
self.__current_run_index = 0
self.__result = "" self.__result = ""
self.__comb_forced_path: list[int] = [] self.__forced_path: list[int] = []
self.__comb_trace: list[int] = [] self.__trace: list[int] = []
self.__cycl_forced_path: list[int] = [] self.__rand_decisions: dict[int, int] = {}
self.__cycl_trace: list[int] = []
def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None): def log(self, kind, message: str, min_level: DEBUG_LEVEL | None = None):
log(self.state.logger, self.state.options.debug_level, kind, message, min_level) log(self.state.logger, self.state.options.debug_level, kind, message, min_level)
def warn_or_stop(self, message: str, e: Exception = None): def warn_or_stop(self, message: str, e: Exception = None):
warn_or_stop(self.state, self.__is_negative, message, e) warn_or_stop(
self.state,
WARN_STOP_WHERE.negative if self.__is_negative else WARN_STOP_WHERE.positive,
message,
e,
)
def __reset_run_state(self): def __reset_run_state(self):
"""Reset all per-run mutable state for a fresh combinatorial pass.""" """Reset all per-run mutable state for a fresh combinatorial/multiple pass."""
self.__shell = [] self.__shell = []
self.__negtags = [] self.__negtags = []
self.__already_processed = [] self.__already_processed = []
@@ -98,6 +106,35 @@ class TreeProcessor(lark.visitors.Interpreter):
if self.state.extranetwork_mappings_obj is not None: if self.state.extranetwork_mappings_obj is not None:
self.state.extranetwork_mappings_obj.cached_mappings.clear() self.state.extranetwork_mappings_obj.cached_mappings.clear()
def __prepare_next_rng(self):
"""
Prepare the random number generator for a new run.
Note: All random generation should use default_rng to be reproducible.
"""
if self.state.options.next_seed == NEXT_SEED.input:
# Use the input seed for every run, so the output is deterministic for a given input.
self.__current_input_seed = (
self.state.inputs.seed
if not isinstance(self.state.inputs.seed, list)
else self.state.inputs.seed[self.__current_run_index % len(self.state.inputs.seed)]
)
self.__rng = np.random.default_rng(self.__current_input_seed)
elif self.state.options.next_seed == NEXT_SEED.increment:
# Increment the seed for each run, so the output varies for a given input.
self.__current_input_seed = clamp_host_bits(self.state.host_config.seed_bits, self.__current_input_seed + 1)
self.__rng = np.random.default_rng(self.__current_input_seed)
elif self.state.options.next_seed == NEXT_SEED.decrement:
# Decrement the seed for each run, so the output varies for a given input.
self.__current_input_seed = clamp_host_bits(self.state.host_config.seed_bits, self.__current_input_seed - 1)
self.__rng = np.random.default_rng(self.__current_input_seed)
elif self.state.options.next_seed == NEXT_SEED.randomize:
# Randomize the seed for each run, so the output varies for a given input.
self.__current_input_seed = self.__rng.integers(
0, 1 << (self.state.host_config.seed_bits - 1), dtype=np.int64
)
self.__rng = np.random.default_rng(self.__current_input_seed)
def start_visit( def start_visit(
self, self,
parsed: lark.Tree, parsed: lark.Tree,
@@ -112,54 +149,89 @@ class TreeProcessor(lark.visitors.Interpreter):
Returns: Returns:
list[tuple[str, list[tuple[str,bool]], dict[str, VariableEntry]]]: A list of list[tuple[str, list[tuple[str,bool]], dict[str, VariableEntry]]]: A list of
(processed prompt, detected wildcards, variables snapshot) triples - one entry per (processed prompt, detected wildcards, variables snapshot) triples - one entry per
combination in combinatorial mode, or a single entry otherwise. The variables snapshot result in combinatorial/multiple mode, or a single entry otherwise. The variables snapshot
is the return value of ``state.variables.backup_user_and_echoed()`` captured after processing. is the return value of ``state.variables.backup_user_and_echoed()`` captured after processing.
""" """
self.log(logging.INFO, "Processing prompt...") self.log(logging.INFO, "Processing prompt...")
self.__detectedWildcards = [] results: list[tuple[str, list[tuple[str, bool]], tuple]] = []
self.__is_negative = False max_results = (
self.__result = "" 1
if self.state.options.run_mode == RUN_MODE.single
or (self.state.options.run_mode == RUN_MODE.multiple and self.state.options.results_limit < 1)
else self.state.options.results_limit
)
self.__current_run_index = 0
if not self.state.options.do_combinatorial: initial_vars = self.state.variables.backup_user()
self.__cycl_forced_path = list(self.state.cyclical_state.current_path)
self.__cycl_trace = [] if self.state.options.run_mode != RUN_MODE.combinatorial:
self.visit(parsed) initial_path = list(self.state.cyclical_state.current_path)
self.__finalize_variables() warned_cycle = False
if self.__cycl_trace: for step in range(max_results):
self.state.cyclical_state.last_trace = self.__cycl_trace[:] self.log(logging.DEBUG, f"Using seed {self.__current_input_seed}")
self.state.cyclical_state.advance() self.__reset_run_state()
return [(self.__result, self.__detectedWildcards, self.state.variables.backup_user())] self.state.variables.restore_user(initial_vars)
self.__forced_path = list(self.state.cyclical_state.current_path)
self.__trace = []
self.visit(parsed)
self.__finalize_variables()
warn = False
if self.__trace:
self.state.cyclical_state.last_trace = self.__trace[:]
self.state.cyclical_state.advance()
if not warned_cycle and (self.state.options.run_mode == RUN_MODE.single or step < max_results - 1):
trace_len = len(self.state.cyclical_state.last_trace)
# Single mode: each start_visit() call advances once, so compare against the
# canonical cycle start (all zeros) to detect a completed global cycle.
# Multiple mode: compare against initial_path to detect the exact repeat point.
if self.state.options.run_mode == RUN_MODE.single:
cycle_start = [0] * trace_len
else:
cycle_start = (initial_path + [0] * trace_len)[:trace_len]
if self.state.cyclical_state.current_path == cycle_start:
warn = True
results.append((self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user_and_output()))
if self.state.options.run_mode == RUN_MODE.multiple:
self.log(logging.INFO, f"Added result {len(results)}")
if warn:
self.log(logging.WARNING, "Cyclical combinations are repeating; prompt results will start repeating.")
warned_cycle = True
self.__current_run_index += 1
self.__prepare_next_rng()
return results
# Combinatorial mode: explore every possible path through choices and wildcards via DFS. # Combinatorial mode: explore every possible path through choices and wildcards via DFS.
# __comb_forced_path drives which option is selected at each decision point; # __forced_path drives which option is selected at each decision point;
# __comb_trace records how many options were available at each point so the DFS can # __trace records how many options were available at each point so the DFS can
# correctly enumerate unexplored branches after each run. # correctly enumerate unexplored branches after each run.
initial_vars = self.state.variables.backup_user() # __rand_decisions caches random (~) choices so they stay consistent across all runs.
results: list[tuple[str, list[tuple[str, bool]], tuple]] = [] self.__rand_decisions = {}
limit = self.state.options.combinatorial_limit
def _run(forced_path: tuple[int, ...]) -> tuple[int, ...]: def _run(forced_path: tuple[int, ...]) -> tuple[int, ...]:
self.log(logging.DEBUG, f"Running combinatorial path: {forced_path}") self.log(logging.DEBUG, f"Running combinatorial path: {forced_path}")
self.__comb_forced_path = list(forced_path) self.log(logging.DEBUG, f"Using seed {self.__current_input_seed}")
self.__comb_trace = [] self.__forced_path = list(forced_path)
self.__trace = []
self.__reset_run_state() self.__reset_run_state()
self.state.variables.restore_user(initial_vars) self.state.variables.restore_user(initial_vars)
self.visit(parsed) self.visit(parsed)
self.__finalize_variables() self.__finalize_variables()
results.append((self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user())) results.append((self.__result, self.__detectedWildcards.copy(), self.state.variables.backup_user_and_output()))
if len(results) == 1: if len(results) == 1:
first_run_estimate = reduce(lambda x, y: x * y, self.__comb_trace, 1) first_run_estimate = reduce(lambda x, y: x * y, self.__trace, 1)
self.log(logging.INFO, f"Estimated combinations (lower bound): {first_run_estimate}") self.log(logging.INFO, f"Estimated combinations (lower bound): {first_run_estimate}")
self.log(logging.INFO, f"Added combination {len(results)}") self.log(logging.INFO, f"Added combination {len(results)}")
return tuple(self.__comb_trace) self.__current_run_index += 1
self.__prepare_next_rng()
return tuple(self.__trace)
limit_reached = False limit_reached = False
def _dfs(forced_path: tuple[int, ...]): def _dfs(forced_path: tuple[int, ...]):
"""Recursively explore combinatorial branches via Depth First Search (DFS).""" """Recursively explore combinatorial branches via Depth First Search (DFS)."""
nonlocal limit_reached nonlocal limit_reached
if 0 < limit <= len(results): if 0 < max_results <= len(results):
limit_reached = True limit_reached = True
return return
trace = _run(forced_path) trace = _run(forced_path)
@@ -168,12 +240,12 @@ class TreeProcessor(lark.visitors.Interpreter):
# Iterating in reverse means later (deeper) decision points vary fastest, # Iterating in reverse means later (deeper) decision points vary fastest,
# so the output order is depth-first rather than breadth-first. # so the output order is depth-first rather than breadth-first.
for i in range(len(trace) - 1, len(forced_path) - 1, -1): for i in range(len(trace) - 1, len(forced_path) - 1, -1):
if 0 < limit <= len(results): if 0 < max_results <= len(results):
limit_reached = True limit_reached = True
return return
num_options = trace[i] num_options = trace[i]
for opt in range(1, num_options): for opt in range(1, num_options):
if 0 < limit <= len(results): if 0 < max_results <= len(results):
limit_reached = True limit_reached = True
return return
# Pad with zeros for intermediate decisions so they keep the default. # Pad with zeros for intermediate decisions so they keep the default.
@@ -182,7 +254,9 @@ class TreeProcessor(lark.visitors.Interpreter):
_dfs(()) _dfs(())
if limit_reached: if limit_reached:
self.log(logging.WARNING, f"Combinatorial limit of {limit} reached; some combinations have been skipped.") self.log(
logging.WARNING, f"Combinatorial limit of {max_results} reached; some combinations have been skipped."
)
return results return results
def __finalize_variables(self): def __finalize_variables(self):
@@ -198,6 +272,8 @@ class TreeProcessor(lark.visitors.Interpreter):
name, specifier = self.__separate_arrayref(k) name, specifier = self.__separate_arrayref(k)
value = self.get_final_scalar_variable(name, specifier) value = self.get_final_scalar_variable(name, specifier)
self.state.variables.set_user(k, value) self.state.variables.set_user(k, value)
# Set the system variable for the current input seed so it can be used afterwards if needed.
self.state.variables.set_system("_output_seed", int(self.__current_input_seed))
def __visit( def __visit(
self, self,
@@ -1288,14 +1364,14 @@ class TreeProcessor(lark.visitors.Interpreter):
f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set." f"Invalid variable name '{escape_single_quotes(variable_name)}' detected! System variables cannot be set."
) )
return return
app = self.state.env_info.get("app", "") app = self.state.env_info.app
if app not in (SUPPORTED_APPS.comfyui.value, SUPPORTED_APPS.tests.value): if app not in (SUPPORTED_APPS.comfyui, SUPPORTED_APPS.tests):
self.warn_or_stop(f"Setting '{escape_single_quotes(variable_name)}' is only supported in ComfyUI.") self.warn_or_stop(f"Setting '{escape_single_quotes(variable_name)}' is only supported in ComfyUI.")
return return
evaluated = self.__visit(content, restore_state=False, discard_content=True) evaluated = self.__visit(content, restore_state=False, discard_content=True)
self.state.env_info[settable_sysvars[variable_name]] = evaluated setattr(self.state.env_info, settable_sysvars[variable_name], evaluated)
if variable_name == "_modelfullname": if variable_name == "_modelfullname":
self.state.env_info["model_class"] = "" # reset model class so it will be re-evaluated self.state.env_info.model_class = "" # reset so it will be re-evaluated
if self.__on_model_info_update is not None: if self.__on_model_info_update is not None:
self.__on_model_info_update() self.__on_model_info_update()
info = variable_name + " = " + f"'{escape_single_quotes(evaluated)}'" info = variable_name + " = " + f"'{escape_single_quotes(evaluated)}'"
@@ -1641,24 +1717,38 @@ class TreeProcessor(lark.visitors.Interpreter):
else_mapping = v else_mapping = v
num_mappings = len(found_mappings) num_mappings = len(found_mappings)
if num_mappings > 0: if num_mappings > 0:
if self.state.options.do_combinatorial: if num_mappings == 1:
decision_idx = len(self.__comb_trace) found = found_mappings[0]
self.__comb_trace.append(num_mappings) elif self.state.options.run_mode == RUN_MODE.combinatorial:
decision_idx = len(self.__trace)
# if self.state.options.default_sampler == DEFAULT_SAMPLER.random:
# # In combinatorial mode with random sampler, optionally fix the choice.
# if decision_idx not in self.__rand_decisions or not self.state.options.comb_random_fixed:
# weights = np.array([float(v.weight or 1) for v in found_mappings])
# weights /= weights.sum()
# self.__rand_decisions[decision_idx] = int(self.__rng.choice(num_mappings, p=weights))
# found = found_mappings[self.__rand_decisions[decision_idx] % num_mappings]
# else: # cyclical: enumerate all variants
self.__trace.append(num_mappings)
chosen_idx = ( chosen_idx = (
min(self.__comb_forced_path[decision_idx], num_mappings - 1) min(self.__forced_path[decision_idx], num_mappings - 1)
if decision_idx < len(self.__comb_forced_path) if decision_idx < len(self.__forced_path)
else 0 else 0
) )
found = found_mappings[chosen_idx] found = found_mappings[chosen_idx]
elif num_mappings == 1: elif self.state.options.default_sampler == DEFAULT_SAMPLER.cyclical:
found = found_mappings[0] decision_idx = len(self.__trace)
else: self.__trace.append(num_mappings)
found = found_mappings[ chosen_idx = (
self.__rng.choice( self.__forced_path[decision_idx] % num_mappings
num_mappings, if decision_idx < len(self.__forced_path)
p=[v.weight or 1 for v in found_mappings], else 0
) )
] found = found_mappings[chosen_idx]
else: # random
weights = np.array([float(v.weight or 1) for v in found_mappings])
weights /= weights.sum()
found = found_mappings[int(self.__rng.choice(num_mappings, p=weights))]
else: else:
found = else_mapping found = else_mapping
# Only cache when at most one mapping matched: with multiple matches, # Only cache when at most one mapping matched: with multiple matches,
@@ -1871,7 +1961,12 @@ class TreeProcessor(lark.visitors.Interpreter):
seen_wildcards_len = len(self.__seen_wildcards) seen_wildcards_len = len(self.__seen_wildcards)
if options is None: if options is None:
options = {} options = {}
sampler: str = options.get("sampler", "~") specified_sampler = options.get("sampler")
sampler: str = (
specified_sampler
if specified_sampler is not None
else "@" if self.state.options.default_sampler == DEFAULT_SAMPLER.cyclical else "~"
)
repeating: bool = options.get("repeating", False) repeating: bool = options.get("repeating", False)
optional: bool = options.get("optional", False) optional: bool = options.get("optional", False)
if "count" in options: if "count" in options:
@@ -1891,8 +1986,12 @@ class TreeProcessor(lark.visitors.Interpreter):
included_choices = 0 included_choices = 0
excluded_choices = 0 excluded_choices = 0
excluded_weights_sum = 0 excluded_weights_sum = 0
else_choice = None
for i, c in enumerate(expanded_choice_values): for i, c in enumerate(expanded_choice_values):
c["choice_index"] = i # we index them to later sort the results c["choice_index"] = i # we index them to later sort the results
if c.get("else", False):
else_choice = c
continue
weight = float(c.get("weight", 1.0)) weight = float(c.get("weight", 1.0))
condition = c.get("if", None) condition = c.get("if", None)
if weight > 0 and (condition is None or self.__eval_condition(condition)): if weight > 0 and (condition is None or self.__eval_condition(condition)):
@@ -1903,10 +2002,15 @@ class TreeProcessor(lark.visitors.Interpreter):
weights.append(-1) weights.append(-1)
excluded_choices += 1 excluded_choices += 1
excluded_weights_sum += weight excluded_weights_sum += weight
if not available_choices and else_choice is not None:
available_choices = [else_choice]
weights = [float(else_choice.get("weight", 1.0))]
included_choices = 1
if excluded_choices > 0: # we need to redistribute the excluded weights if excluded_choices > 0: # we need to redistribute the excluded weights
weights = [weight + excluded_weights_sum / included_choices for weight in weights if weight >= 0] weights = [weight + excluded_weights_sum / included_choices for weight in weights if weight >= 0]
weights = np.array(weights) weights = np.array(weights)
weights /= weights.sum() # normalize weights weights /= weights.sum() # normalize weights
selected_choices: list[dict] = []
if available_choices: if available_choices:
if from_value < 0: if from_value < 0:
from_value = 1 from_value = 1
@@ -1916,8 +2020,7 @@ class TreeProcessor(lark.visitors.Interpreter):
to_value = 1 to_value = 1
elif (to_value > len(available_choices) and not repeating) or from_value > to_value: elif (to_value > len(available_choices) and not repeating) or from_value > to_value:
to_value = len(available_choices) to_value = len(available_choices)
comb_chosen_selection: Optional[list[dict]] = None if self.state.options.run_mode == RUN_MODE.combinatorial or sampler == "@":
if self.state.options.do_combinatorial or sampler == "@":
# Enumerate every distinct selection of choices, accounting for count range and repetition. # Enumerate every distinct selection of choices, accounting for count range and repetition.
all_selections: list[tuple] = [] all_selections: list[tuple] = []
# When keep_choices_order is False the output depends on the selection order, # When keep_choices_order is False the output depends on the selection order,
@@ -1936,71 +2039,74 @@ class TreeProcessor(lark.visitors.Interpreter):
all_selections.extend(combinations(available_choices, k)) all_selections.extend(combinations(available_choices, k))
else: else:
all_selections.extend(permutations(available_choices, k)) all_selections.extend(permutations(available_choices, k))
num_selections = len(all_selections) decision_idx = len(self.__trace)
if self.state.options.do_combinatorial: if self.state.options.run_mode == RUN_MODE.combinatorial:
decision_idx = len(self.__comb_trace) if specified_sampler == "~":
self.__comb_trace.append(num_selections) # In combinatorial mode with a explicit random sampler, fix the random choice
# once across all DFS runs so every combination uses the same value.
if decision_idx not in self.__rand_decisions or not self.state.options.comb_random_fixed:
# We need to make a random choice for this decision index and store it for future runs.
# If comb_random_fixed is True, we only do this once per decision index, so all
# combinations share the same choice.
# If False, we do it every time, which allows for different random choices across DFS runs.
self.__rand_decisions[decision_idx] = int(self.__rng.choice(len(all_selections)))
all_selections = [all_selections[self.__rand_decisions[decision_idx] % len(all_selections)]]
num_selections = len(all_selections)
self.__trace.append(num_selections)
chosen_idx = ( chosen_idx = (
min(self.__comb_forced_path[decision_idx], num_selections - 1) min(self.__forced_path[decision_idx], num_selections - 1)
if decision_idx < len(self.__comb_forced_path) if decision_idx < len(self.__forced_path)
else 0 else 0
) )
else: # sampler == "@" else: # sampler == "@"
cycl_decision_idx = len(self.__cycl_trace) num_selections = len(all_selections)
self.__cycl_trace.append(num_selections) self.__trace.append(num_selections)
chosen_idx = ( chosen_idx = (
self.__cycl_forced_path[cycl_decision_idx] % num_selections self.__forced_path[decision_idx] % num_selections
if cycl_decision_idx < len(self.__cycl_forced_path) if decision_idx < len(self.__forced_path)
else 0 else 0
) )
comb_chosen_selection = list(all_selections[chosen_idx]) selected_choices = list(all_selections[chosen_idx])
num_choices = len(comb_chosen_selection)
else: else:
num_choices = ( chosen_num_choices = (
self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value self.__rng.integers(from_value, to_value, endpoint=True) if from_value < to_value else from_value
) )
if chosen_num_choices < 2:
repeating = False
selected_choices = (
list(self.__rng.choice(available_choices, size=chosen_num_choices, p=weights, replace=repeating))
if available_choices
else []
)
else: else:
num_choices = 0
if not optional and from_value > 0: if not optional and from_value > 0:
self.warn_or_stop(f"Not enough choices found for {msg_where}!") self.warn_or_stop(f"Not enough choices found for {msg_where}!")
if num_choices < 2: if self.state.options.keep_choices_order:
repeating = False selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"])
num_choices = len(selected_choices)
self.log( self.log(
logging.DEBUG, logging.DEBUG,
f"Selecting {'optional ' if optional else ''}{'repeating ' if repeating else ''}{num_choices} choice" f"Selecting {'optional ' if optional else ''}{'repeating ' if repeating else ''}{num_choices} choice"
+ ("s" if num_choices != 1 else "") + ("s" if num_choices != 1 else "")
+ (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else ""), + (f" and separating with '{escape_single_quotes(separator)}'" if num_choices > 1 else ""),
) )
if num_choices > 0: selected_choices_text = []
if comb_chosen_selection is not None: for i, c in enumerate(selected_choices):
selected_choices: list[dict] = comb_chosen_selection t1 = time.monotonic_ns()
choice_content_obj = c.get("content", c.get("text", None))
if isinstance(choice_content_obj, str):
choice_content = choice_content_obj
else: else:
selected_choices: list[dict] = ( choice_content = self.__visit(choice_content_obj, False, True)
list(self.__rng.choice(available_choices, size=num_choices, p=weights, replace=repeating)) t2 = time.monotonic_ns()
if available_choices self.log(
else [] logging.DEBUG,
) f"Adding choice {i+1} ({(t2-t1) / 1_000_000_000:.3f} seconds):\n"
if self.state.options.keep_choices_order: + textwrap.indent(re.sub(r"\n$", "", choice_content), " "),
selected_choices = sorted(selected_choices, key=lambda x: x["choice_index"]) )
selected_choices_text = [] selected_choices_text.append(choice_content)
for i, c in enumerate(selected_choices): # remove comments
t1 = time.monotonic_ns() results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text]
choice_content_obj = c.get("content", c.get("text", None))
if isinstance(choice_content_obj, str):
choice_content = choice_content_obj
else:
choice_content = self.__visit(choice_content_obj, False, True)
t2 = time.monotonic_ns()
self.log(
logging.DEBUG,
f"Adding choice {i+1} ({(t2-t1) / 1_000_000_000:.3f} seconds):\n"
+ textwrap.indent(re.sub(r"\n$", "", choice_content), " "),
)
selected_choices_text.append(choice_content)
# remove comments
results = [re.sub(r"\s*#[^\n]*(?:\n|$)", "", r, flags=re.DOTALL) for r in selected_choices_text]
else:
results = []
container = options.get("container", None) container = options.get("container", None)
if container is None: if container is None:
separator = options.get("separator", self.state.options.choice_separator) separator = options.get("separator", self.state.options.choice_separator)
@@ -2030,11 +2136,12 @@ class TreeProcessor(lark.visitors.Interpreter):
), ),
], ],
) )
self.log( if self.__seen_wildcards[seen_wildcards_len:]:
logging.DEBUG, self.log(
"Unseen wildcards: " logging.DEBUG,
+ ", ".join([f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]), "Unseen wildcards: "
) + ", ".join([f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]),
)
self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len] self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len]
return container, results return container, results
@@ -2114,7 +2221,12 @@ class TreeProcessor(lark.visitors.Interpreter):
c_label_obj = choice.children[1] c_label_obj = choice.children[1]
choice_dict["labels"] = [str(x).lower() for x in c_label_obj.children[1:-1]] if c_label_obj is not None else [] choice_dict["labels"] = [str(x).lower() for x in c_label_obj.children[1:-1]] if c_label_obj is not None else []
choice_dict["weight"] = float(choice.children[2].children[0]) if choice.children[2] is not None else 1.0 choice_dict["weight"] = float(choice.children[2].children[0]) if choice.children[2] is not None else 1.0
choice_dict["if"] = choice.children[3].children[0] if choice.children[3] is not None else None ifelse = choice.children[3]
if ifelse is not None:
if len(ifelse.children) and ifelse.children[0] is not None:
choice_dict["if"] = ifelse.children[0]
else:
choice_dict["else"] = True
choice_dict["content"] = choice.children[-1] choice_dict["content"] = choice.children[-1]
return choice_dict return choice_dict
@@ -2242,7 +2354,7 @@ class TreeProcessor(lark.visitors.Interpreter):
parse_prompt( parse_prompt(
self.state, self.state,
"as wildcard options", "as wildcard options",
wildcard.unprocessed_choices[0][:-2].strip(), wildcard.unprocessed_choices[0],
self.state.parsers["wcdefoptions"], self.state.parsers["wcdefoptions"],
True, True,
), ),
@@ -2371,7 +2483,8 @@ class TreeProcessor(lark.visitors.Interpreter):
self.__result += wc self.__result += wc
if self.__debug_level == DEBUG_LEVEL.full: if self.__debug_level == DEBUG_LEVEL.full:
list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]] list_unseen = [f"'{escape_single_quotes(x)}'" for x in self.__seen_wildcards[seen_wildcards_len:]]
self.log(logging.DEBUG, f"Unseen wildcards: {', '.join(list_unseen)}") if list_unseen:
self.log(logging.DEBUG, f"Unseen wildcards: {', '.join(list_unseen)}")
self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len] self.__seen_wildcards = self.__seen_wildcards[:seen_wildcards_len]
t2 = time.monotonic_ns() t2 = time.monotonic_ns()
self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'") self.__debug_end("wildcard", start_result, t2 - t1, f"'{escape_single_quotes(wc)}'")
+7
View File
@@ -138,6 +138,13 @@ class VariableRepository:
} }
) )
def backup_user_and_output(self) -> dict[str, VariableEntry]:
"""Return a per-entry shallow-copy snapshot of all user variables and output variables."""
return {
name: VariableEntry(entry.value, entry.last_echoed_value, entry.last_echoed_evaluated_value)
for name, entry in self._vars.items()
} | {name: VariableEntry(entry) for name, entry in self._system.items() if name.startswith("_output_")}
# ---- Combined queries ---- # ---- Combined queries ----
def get(self, name: str, default: Any = None) -> Any: def get(self, name: str, default: Any = None) -> Any:
+46 -9
View File
@@ -1,5 +1,6 @@
import fnmatch import fnmatch
from pathlib import Path from pathlib import Path
import re
from typing import Any, Optional from typing import Any, Optional
import logging import logging
from ruamel.yaml import YAML as _YAML from ruamel.yaml import YAML as _YAML
@@ -114,23 +115,33 @@ class PPPWildcards:
keys = sorted(fnmatch.filter(self.wildcards.keys(), key)) keys = sorted(fnmatch.filter(self.wildcards.keys(), key))
return [self.wildcards[k] for k in keys] return [self.wildcards[k] for k in keys]
def __get_wc_in_dict(self, dictionary: dict, prefix="") -> list[tuple[str, Any]]: def __get_wc_in_dict(self, dictionary: dict, prefix="", file_str: str = "") -> list[tuple[str, Any]]:
""" """
Get all wildcards in a dictionary, along their object. Get all wildcards in a dictionary, along their object.
Args: Args:
dictionary (dict): The dictionary to check. dictionary (dict): The dictionary to check.
prefix (str): The prefix for the current key. prefix (str): The prefix for the current key.
file_str (str): The file string for logging purposes.
Returns: Returns:
list: A list of all leaf wildcards in the dictionary. list: A list of all leaf wildcards in the dictionary.
""" """
wc = [] wc = []
for key, obj in dictionary.items(): for key, obj in dictionary.items():
strkey = str(key)
if not self.__check_key_validity(strkey):
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard name part '{escape_single_quotes(prefix + strkey)}' in file '{escape_single_quotes(file_str)}'!",
)
continue
if isinstance(obj, dict): if isinstance(obj, dict):
wc.extend(self.__get_wc_in_dict(obj, prefix + str(key) + "/")) wc.extend(self.__get_wc_in_dict(obj, prefix + strkey + "/", file_str))
else: else:
wc.append((prefix + str(key), obj)) wc.append((prefix + strkey, obj))
return wc return wc
def __remove_wildcards_from_path(self, full_path: Path, debug=True): def __remove_wildcards_from_path(self, full_path: Path, debug=True):
@@ -218,7 +229,7 @@ class PPPWildcards:
wildcards_input = wildcards_input.strip() wildcards_input = wildcards_input.strip()
if wildcards_input != "": if wildcards_input != "":
try: try:
content = _YAML(typ='safe').load(wildcards_input) content = _YAML(typ="safe").load(wildcards_input)
except _YAMLError as e: except _YAMLError as e:
log(self.__logger, self.__debug_level, logging.WARNING, f"Invalid format for input wildcards: {e}") log(self.__logger, self.__debug_level, logging.WARNING, f"Invalid format for input wildcards: {e}")
return return
@@ -268,7 +279,7 @@ class PPPWildcards:
Returns: Returns:
bool: Whether the dictionary is a valid choice options dictionary or not. bool: Whether the dictionary is a valid choice options dictionary or not.
""" """
return all(k in ["command", "labels", "weight", "if", "content", "text"] for k in d.keys()) return all(k in ["command", "labels", "weight", "if", "else", "content", "text"] for k in d.keys())
def __get_choices(self, obj: object, full_path: Path | None, key_parts: list[str]) -> list: def __get_choices(self, obj: object, full_path: Path | None, key_parts: list[str]) -> list:
""" """
@@ -374,6 +385,21 @@ class PPPWildcards:
value = f"{options}::{value}" value = f"{options}::{value}"
return value return value
def __check_key_validity(self, key: str) -> bool:
"""
Check if a key is valid.
Args:
key (str): The key to check.
Returns:
bool: Whether the key is valid or not.
"""
match = re.match(r"^[a-zA-Z0-9\-_#]+$", key)
if match is None:
return False
return True
def __add_wildcard(self, content: object, full_path: Path | None, external_key_parts: list[str]): def __add_wildcard(self, content: object, full_path: Path | None, external_key_parts: list[str]):
""" """
Add a wildcard to the wildcards dictionary. Add a wildcard to the wildcards dictionary.
@@ -389,9 +415,20 @@ class PPPWildcards:
return str(wc.file) if wc.file is not None else "input" return str(wc.file) if wc.file is not None else "input"
key_parts = external_key_parts.copy() key_parts = external_key_parts.copy()
if key_parts:
for part in key_parts:
strkey = str(part)
if strkey and not self.__check_key_validity(strkey):
log(
self.__logger,
self.__debug_level,
logging.WARNING,
f"Invalid wildcard name start '{escape_single_quotes('/'.join(key_parts))}' in file '{escape_single_quotes(file_str)}'!",
)
return
if isinstance(content, dict): if isinstance(content, dict):
key_parts.pop() key_parts.pop() # we don't want the name of the filename to be included, since the dict keys will be used instead
keys = self.__get_wc_in_dict(content) keys = self.__get_wc_in_dict(content, "", file_str)
for key, obj in keys: for key, obj in keys:
tmp_key_parts = key_parts.copy() tmp_key_parts = key_parts.copy()
tmp_key_parts.extend(key.split("/")) tmp_key_parts.extend(key.split("/"))
@@ -472,7 +509,7 @@ class PPPWildcards:
external_key_parts = list(full_path.with_suffix("").relative_to(base).parts) external_key_parts = list(full_path.with_suffix("").relative_to(base).parts)
try: try:
with open(full_path, "r", encoding="utf-8") as file: with open(full_path, "r", encoding="utf-8") as file:
content = _YAML(typ='safe').load(file) content = _YAML(typ="safe").load(file)
except: # pylint: disable=bare-except except: # pylint: disable=bare-except
log( log(
self.__logger, self.__logger,
@@ -481,7 +518,7 @@ class PPPWildcards:
f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...", f"Could not read file '{escape_single_quotes(str(full_path))}' with utf-8 encoding, trying windows-1252...",
) )
with open(full_path, "r", encoding="windows-1252") as file: with open(full_path, "r", encoding="windows-1252") as file:
content = _YAML(typ='safe').load(file) content = _YAML(typ="safe").load(file)
self.__add_wildcard(content, full_path, external_key_parts) self.__add_wildcard(content, full_path, external_key_parts)
def __get_wildcards_in_text_file(self, full_path: Path, base: Path): def __get_wildcards_in_text_file(self, full_path: Path, base: Path):
+8 -7
View File
@@ -1,18 +1,19 @@
[project] [project]
name = "sd-webui-prompt-postprocessor" name = "sd-webui-prompt-postprocessor"
description = "Stable Diffusion WebUI & ComfyUI extension to post-process the prompt. Features include: wildcards, sending content from the prompt to the negative prompt, variables, model detection, extranetwork mapping, cleanup." description = "Stable Diffusion WebUI & ComfyUI extension to post-process the prompt. Features include: wildcards, sending content from the prompt to the negative prompt, variables, model detection, extranetwork mapping, cleanup."
version = "3.2.1" version = "3.3.0"
license = { file = "LICENSE.txt" } license = { file = "LICENSE.txt" }
dependencies = ["lark", "numpy", "ruamel.yaml", "pydantic"] dependencies = ["lark==1.*", "numpy==2.*", "ruamel.yaml==0.*", "pydantic==2.*"]
requires-python = ">=3.10" requires-python = ">=3.10"
[project.urls] [project.urls]
repository = "https://github.com/acorderob/sd-webui-prompt-postprocessor" Repository = "https://github.com/acorderob/sd-webui-prompt-postprocessor"
# Used by Comfy Registry https://registry.comfy.org/ # Used by Comfy Registry https://registry.comfy.org/
documentation = "https://github.com/acorderob/sd-webui-prompt-postprocessor/main/README.md" Documentation = "https://github.com/acorderob/sd-webui-prompt-postprocessor/main/README.md"
"Bug Tracker" = "https://github.com/acorderob/sd-webui-prompt-postprocessor/issues"
issues = "https://github.com/acorderob/sd-webui-prompt-postprocessor/issues" issues = "https://github.com/acorderob/sd-webui-prompt-postprocessor/issues"
[tool.comfy] [tool.comfy]
publisher_id = "acorderob" PublisherId = "acorderob"
display_name = "ACB Prompt PostProcessor" DisplayName = "ACB Prompt PostProcessor"
icon = "https://raw.githubusercontent.com/acorderob/sd-webui-prompt-postprocessor/main/images/prompt-postprocessor-icon.png" Icon = "https://raw.githubusercontent.com/acorderob/sd-webui-prompt-postprocessor/main/images/prompt-postprocessor-icon.png"
+4 -4
View File
@@ -1,4 +1,4 @@
lark lark==1.*
numpy numpy==2.*
ruamel.yaml ruamel.yaml==0.*
pydantic pydantic==2.*
+214 -122
View File
@@ -1,6 +1,7 @@
if __name__ == "__main__": if __name__ == "__main__":
raise SystemExit("This script must be run from a Stable Diffusion WebUI") raise SystemExit("This script must be run from a Stable Diffusion WebUI")
# pylint: disable=wrong-import-position,wrong-import-order
import logging import logging
import sys import sys
import os import os
@@ -10,14 +11,24 @@ import numpy as np
sys.path.append(str(Path(__file__).parent)) # base path for the extension sys.path.append(str(Path(__file__).parent)) # base path for the extension
from modules import scripts, shared, script_callbacks # type: ignore from modules import scripts, shared, script_callbacks # type: ignore # pylint: disable=import-error
from modules.processing import StableDiffusionProcessing # type: ignore from modules.processing import StableDiffusionProcessing # type: ignore # pylint: disable=import-error
from modules.shared import opts # type: ignore from modules.shared import opts # type: ignore # pylint: disable=import-error
from modules.paths import models_path # type: ignore from modules.paths import models_path # type: ignore # pylint: disable=import-error
import gradio as gr # type: ignore import gradio as gr # type: ignore # pylint: disable=import-error
from ppp import PromptPostProcessor from ppp import PromptPostProcessor
from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, SUPPORTED_APPS, SUPPORTED_APPS_NAMES, PPPStateOptions from ppp_classes import (
DEFAULT_SAMPLER,
IFWILDCARDS_CHOICES,
ONWARNING_CHOICES,
PPPEnvInfo,
SUPPORTED_APPS,
SUPPORTED_APPS_NAMES,
RUN_MODE,
NEXT_SEED,
PPPStateOptions,
)
from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log from ppp_logging import DEBUG_LEVEL, PromptPostProcessorLogFactory, log
from ppp_cache import PPPLRUCache from ppp_cache import PPPLRUCache
from ppp_wildcards import PPPWildcards from ppp_wildcards import PPPWildcards
@@ -74,11 +85,13 @@ class PromptPostProcessorA1111Script(scripts.Script):
self.wildcards_obj = None self.wildcards_obj = None
self.extranetwork_mappings_obj = None self.extranetwork_mappings_obj = None
self.ppp_init = False self.ppp_init = False
self.ppp = None
# log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.INFO, f"Initializing {self.name} instance {self.instance_index}") # log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.INFO, f"Initializing {self.name} instance {self.instance_index}")
try: try:
# Support for SD.Next # Support for SD.Next
import installer # type: ignore import installer # type: ignore # pylint: disable=import-outside-toplevel
if hasattr(installer, "control_extensions"): if hasattr(installer, "control_extensions"):
if self.title() not in installer.control_extensions: if self.title() not in installer.control_extensions:
installer.control_extensions.append(self.title()) # We add the extension to the whitelist. installer.control_extensions.append(self.title()) # We add the extension to the whitelist.
@@ -86,7 +99,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.") # log(self.ppp_logger, DEBUG_LEVEL.minimal, logging.WARNING, "Could not import control_extensions from installer, SD.Next support will not work.")
pass pass
except Exception as e: # pylint: disable=broad-except 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): def title(self):
""" """
@@ -113,7 +131,7 @@ class PromptPostProcessorA1111Script(scripts.Script):
with gr.Accordion(PromptPostProcessor.NAME, open=False): with gr.Accordion(PromptPostProcessor.NAME, open=False):
force_equal_seeds = gr.Checkbox( force_equal_seeds = gr.Checkbox(
label="Force equal seeds", label="Force equal seeds",
info="Force all image seeds and variation seeds to be equal to the first one, disabling the default autoincrease.", info="Force all image seeds and variation seeds to be equal to the first one, disabling the default autoincrement.",
value=False, value=False,
# show_label=True, # show_label=True,
elem_id="ppp_force_equal_seeds", elem_id="ppp_force_equal_seeds",
@@ -126,8 +144,6 @@ class PromptPostProcessorA1111Script(scripts.Script):
* A seed of -1 and "Incremental seed" unchecked will use a random seed for each prompt. * 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" 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. * 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.
""") """)
gr.HTML("<br>") gr.HTML("<br>")
with gr.Row(equal_height=True): with gr.Row(equal_height=True):
@@ -155,34 +171,51 @@ class PromptPostProcessorA1111Script(scripts.Script):
elem_id="ppp_incremental_seed", elem_id="ppp_incremental_seed",
) )
gr.HTML("<br>") gr.HTML("<br>")
run_mode = gr.Radio(
label="Run mode",
choices=[rm.value for rm in RUN_MODE],
value=PromptPostProcessor.DEFAULT_RUN_MODE,
info="Select the run mode for prompt processing. 'Single' produces one result, 'Multiple' produces many results to fill the batch, and 'Combinatorial' generates all prompt combinations and cycles through them to fill the batch.",
elem_id="ppp_run_mode",
)
with gr.Row(equal_height=True): with gr.Row(equal_height=True):
combinatorial = gr.Checkbox( results_limit = gr.Number(
label="Combinatorial mode", label="Results limit (0 = no limit)",
info="Generate all prompt combinations and cycle through them to fill the batch.", value=PromptPostProcessor.DEFAULT_RESULTS_LIMIT,
value=PromptPostProcessor.DEFAULT_DO_COMBINATORIAL,
elem_id="ppp_combinatorial",
)
combinatorial_shuffle = gr.Checkbox(
label="Shuffle combinations",
info="Shuffle the combinatorial results.",
value=PromptPostProcessor.DEFAULT_COMBINATORIAL_SHUFFLE,
elem_id="ppp_combinatorial_shuffle",
)
combinatorial_limit = gr.Number(
label="Combinations limit (0 = no limit)",
value=PromptPostProcessor.DEFAULT_COMBINATORIAL_LIMIT,
precision=0, precision=0,
min_width=120, min_width=120,
elem_id="ppp_combinatorial_limit", elem_id="ppp_results_limit",
)
default_sampler = gr.Radio(
label="Default sampler",
choices=[ds.value for ds in DEFAULT_SAMPLER],
value=PromptPostProcessor.DEFAULT_DEFAULT_SAMPLER,
info="Select the default sampler.",
elem_id="ppp_default_sampler",
)
with gr.Row(equal_height=True):
results_shuffle = gr.Checkbox(
label="Shuffle results",
info="Shuffle the results.",
value=PromptPostProcessor.DEFAULT_RESULTS_SHUFFLE,
elem_id="ppp_results_shuffle",
)
comb_random_fixed = gr.Checkbox(
label="Fix random sampler across combinations",
info="Fix the value of any specified random samplers across all combinations in combinatorial mode.",
value=PromptPostProcessor.DEFAULT_COMB_RANDOM_FIXED,
elem_id="ppp_comb_random_fixed",
) )
return [ return [
force_equal_seeds, force_equal_seeds,
unlink_seed, unlink_seed,
seed, seed,
incremental_seed, incremental_seed,
combinatorial, run_mode,
combinatorial_shuffle, results_limit,
combinatorial_limit, results_shuffle,
comb_random_fixed,
default_sampler,
] ]
def process( def process(
@@ -192,9 +225,11 @@ class PromptPostProcessorA1111Script(scripts.Script):
input_unlink_seed, input_unlink_seed,
input_seed, input_seed,
input_incremental_seed, input_incremental_seed,
input_combinatorial, input_run_mode,
input_combinatorial_shuffle, input_results_limit,
input_combinatorial_limit, input_results_shuffle,
input_comb_random_fixed,
input_default_sampler,
): # pylint: disable=arguments-differ ): # pylint: disable=arguments-differ
""" """
Processes the prompts and applies post-processing operations. Processes the prompts and applies post-processing operations.
@@ -205,9 +240,11 @@ class PromptPostProcessorA1111Script(scripts.Script):
input_unlink_seed (bool): Flag indicating whether to unlink the seed. input_unlink_seed (bool): Flag indicating whether to unlink the seed.
input_seed (int): The seed value. input_seed (int): The seed value.
input_incremental_seed (bool): Flag indicating whether to use incremental seed. input_incremental_seed (bool): Flag indicating whether to use incremental seed.
input_combinatorial (bool): Flag indicating whether to use combinatorial mode. input_run_mode (str): The run mode for prompt processing.
input_combinatorial_shuffle (bool): Flag indicating whether to shuffle the combinatorial results. input_results_limit (int): Maximum number of results (0 = no limit).
input_combinatorial_limit (int): Maximum number of combinations (0 = no limit). input_results_shuffle (bool): Flag indicating whether to shuffle the results.
input_comb_random_fixed (bool): Flag indicating whether to fix the random sampler across all combinations.
input_default_sampler (str): The default sampler for prompt processing.
Returns: Returns:
None None
@@ -267,10 +304,15 @@ class PromptPostProcessorA1111Script(scripts.Script):
cup_remove_extranetwork_tags=getattr( cup_remove_extranetwork_tags=getattr(
opts, "ppp_rem_removeextranetworktags", PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS opts, "ppp_rem_removeextranetworktags", PromptPostProcessor.DEFAULT_CUP_REMOVE_EXTRANETWORK_TAGS
), ),
do_combinatorial=input_combinatorial,
combinatorial_shuffle=input_combinatorial_shuffle,
combinatorial_limit=max(num_seeds, int(input_combinatorial_limit)) if input_combinatorial else 0,
results_file=getattr(opts, "ppp_gen_resultsfile", PromptPostProcessor.DEFAULT_RESULTS_FILE), results_file=getattr(opts, "ppp_gen_resultsfile", PromptPostProcessor.DEFAULT_RESULTS_FILE),
run_mode=RUN_MODE(input_run_mode if input_run_mode else PromptPostProcessor.DEFAULT_RUN_MODE),
results_limit=min(num_seeds, int(input_results_limit)),
results_shuffle=input_results_shuffle,
comb_random_fixed=input_comb_random_fixed,
default_sampler=DEFAULT_SAMPLER(
input_default_sampler if input_default_sampler else PromptPostProcessor.DEFAULT_DEFAULT_SAMPLER
),
next_seed=NEXT_SEED.input, # we use the calculated seeds
) )
if not self.ppp_init: if not self.ppp_init:
self.ppp_init = True self.ppp_init = True
@@ -301,9 +343,16 @@ class PromptPostProcessorA1111Script(scripts.Script):
"PPP unlink seed": input_unlink_seed, "PPP unlink seed": input_unlink_seed,
"PPP prompt seed": input_seed, "PPP prompt seed": input_seed,
"PPP incremental seed": input_incremental_seed, "PPP incremental seed": input_incremental_seed,
"PPP combinatorial": input_combinatorial, "PPP run mode": input_run_mode,
"PPP default sampler": input_default_sampler,
} }
) )
if input_run_mode == RUN_MODE.combinatorial.value:
p.extra_generation_params.update(
{
"PPP combinatorial random fixed": input_comb_random_fixed,
}
)
log( log(
self.ppp_logger, self.ppp_logger,
@@ -311,17 +360,17 @@ class PromptPostProcessorA1111Script(scripts.Script):
logging.INFO, logging.INFO,
f"Post-processing prompts ({'i2i' if is_i2i else 't2i'})", f"Post-processing prompts ({'i2i' if is_i2i else 't2i'})",
) )
env_info = { env_info = PPPEnvInfo(
"app": app.value, app=app,
"models_path": models_path, models_path=models_path,
"model_filename": getattr(p.sd_model.sd_checkpoint_info, "filename", ""), model_filename=getattr(p.sd_model.sd_checkpoint_info, "filename", ""),
"model_class": ( model_class=(
p.sd_model.model_config.__class__.__name__ p.sd_model.model_config.__class__.__name__
if app in (SUPPORTED_APPS.forge, SUPPORTED_APPS.forgeneo) if app in (SUPPORTED_APPS.forge, SUPPORTED_APPS.forgeneo)
else p.sd_model.__class__.__name__ else p.sd_model.__class__.__name__
), ),
"property_base": p.sd_model, property_base=p.sd_model,
} )
wc_wildcards_folders = getattr(opts, "ppp_wil_wildcardsfolders", "") wc_wildcards_folders = getattr(opts, "ppp_wil_wildcardsfolders", "")
if wc_wildcards_folders == "": if wc_wildcards_folders == "":
wc_wildcards_folders = os.getenv("WILDCARD_DIR", PPPWildcards.DEFAULT_WILDCARDS_FOLDER) wc_wildcards_folders = os.getenv("WILDCARD_DIR", PPPWildcards.DEFAULT_WILDCARDS_FOLDER)
@@ -345,16 +394,26 @@ class PromptPostProcessorA1111Script(scripts.Script):
self.ppp_debug_level, wildcards_folders if options.process_wildcards else None self.ppp_debug_level, wildcards_folders if options.process_wildcards else None
) )
self.extranetwork_mappings_obj.refresh_extranetwork_mappings(self.ppp_debug_level, enmappings_folders) self.extranetwork_mappings_obj.refresh_extranetwork_mappings(self.ppp_debug_level, enmappings_folders)
ppp = PromptPostProcessor( if self.ppp is None:
self.ppp_logger, self.ppp = PromptPostProcessor(
env_info, self.ppp_logger,
options, env_info,
self.grammar_content, options,
self.ppp_interrupt, self.grammar_content,
self.wildcards_obj, self.ppp_interrupt,
self.extranetwork_mappings_obj, self.wildcards_obj,
self.extranetwork_mappings_obj,
)
else:
self.ppp.update(
env_info,
options,
self.wildcards_obj,
self.extranetwork_mappings_obj,
)
hash_fullenv = hash(
(self.ppp.envinfo_hash, self.ppp.options_hash, self.wildcards_obj, self.extranetwork_mappings_obj)
) )
hash_fullenv = hash((ppp.envinfo_hash, ppp.options_hash, self.wildcards_obj, self.extranetwork_mappings_obj))
if input_force_equal_seeds: if input_force_equal_seeds:
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Forcing equal seeds") log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Forcing equal seeds")
@@ -367,10 +426,10 @@ class PromptPostProcessorA1111Script(scripts.Script):
if input_unlink_seed: if input_unlink_seed:
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Using unlinked seed") log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "Using unlinked seed")
if input_incremental_seed: if input_incremental_seed:
first_seed = np.random.randint(0, 2**32, dtype=np.int64) if input_seed == -1 else input_seed first_seed = np.random.randint(0, 1 << self.ppp.state.host_config.seed_bits, dtype=np.int64) if input_seed == -1 else input_seed
calculated_seeds = [first_seed + i for i in range(num_seeds)] calculated_seeds = [first_seed + i for i in range(num_seeds)]
elif input_seed == -1: elif input_seed == -1:
calculated_seeds = np.random.randint(0, 2**32, size=num_seeds, dtype=np.int64) calculated_seeds = np.random.randint(0, 1 << self.ppp.state.host_config.seed_bits, size=num_seeds, dtype=np.int64)
else: else:
calculated_seeds = [input_seed for _ in range(num_seeds)] calculated_seeds = [input_seed for _ in range(num_seeds)]
else: else:
@@ -388,8 +447,8 @@ class PromptPostProcessorA1111Script(scripts.Script):
else: else:
calculated_seeds = seeds calculated_seeds = seeds
# (prompt type, typeindex) -> (new positive prompt, new negative prompt) # [index][indextype] -> (new positive prompt, new negative prompt)
prompts_list: dict[tuple[str, int], tuple[str, str]] = {} prompts_list: list[list[tuple[str, str]]] = []
extra_params = {} extra_params = {}
# adds prompts # adds prompts
@@ -402,24 +461,27 @@ class PromptPostProcessorA1111Script(scripts.Script):
rnh: list[str] = getattr(p, "all_hr_negative_prompts", None) rnh: list[str] = getattr(p, "all_hr_negative_prompts", None)
hiresfix_exists = bool(rph) and bool(rnh) hiresfix_exists = bool(rph) and bool(rnh)
for i in range(len(calculated_seeds)): for i in range(len(calculated_seeds)):
if regular_exists: prompts_list.append([])
prompts_list[(regular_type, i)] = None prompts_list[i].append(None)
if hiresfix_exists: if hiresfix_exists:
prompts_list[(hiresfix_type, i)] = None prompts_list[i].append(None)
ppp.process_prompts_group_start() self.ppp.process_prompts_group_start()
if input_combinatorial: if input_run_mode in (RUN_MODE.multiple.value, RUN_MODE.combinatorial.value):
seed_for_comb = calculated_seeds[0] if calculated_seeds else 0
regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None) regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None)
hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None) hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None)
regular_changes = False regular_changes = False
hiresfix_changes = False hiresfix_changes = False
if regular_exists: if regular_exists:
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "processing prompts combinatorially (regular)") if input_run_mode == RUN_MODE.combinatorial.value:
comb_results = ppp.process_prompt( msg = "processing prompts combinatorially (regular)"
else:
msg = "processing prompts for multiple results (regular)"
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, msg)
comb_results = self.ppp.process_prompt(
rpr[0], rpr[0],
rnr[0], rnr[0],
seed_for_comb, calculated_seeds,
jobinfo={ jobinfo={
"job_timestamp": shared.state.job_timestamp, "job_timestamp": shared.state.job_timestamp,
"job": shared.state.job, "job": shared.state.job,
@@ -429,8 +491,12 @@ class PromptPostProcessorA1111Script(scripts.Script):
num_comb = len(comb_results) num_comb = len(comb_results)
for i in range(len(rpr)): # pylint: disable=consider-using-enumerate for i in range(len(rpr)): # pylint: disable=consider-using-enumerate
posp, negp, _ = comb_results[i % num_comb] 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 input_run_mode == RUN_MODE.combinatorial.value:
field_name = "PPP combination"
else:
field_name = "PPP result"
extra_params[field_name] = [str(1 + (i % num_comb)) for i in range(len(rpr))]
if hiresfix_exists: if hiresfix_exists:
hiresfix_equal = regular_exists and rph == rpr and rnh == rnr hiresfix_equal = regular_exists and rph == rpr and rnh == rnr
if hiresfix_equal: if hiresfix_equal:
@@ -438,21 +504,25 @@ class PromptPostProcessorA1111Script(scripts.Script):
self.ppp_logger, self.ppp_logger,
self.ppp_debug_level, self.ppp_debug_level,
logging.INFO, logging.INFO,
"hiresfix prompts are the same as regular prompts, skipping combinatorial processing for hiresfix", "hiresfix prompts are the same as regular prompts, skipping processing for hiresfix",
) )
for i in range(len(rph)): # pylint: disable=consider-using-enumerate 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: else:
if input_run_mode == RUN_MODE.combinatorial.value:
msg = "processing prompts combinatorially (hiresfix)"
else:
msg = "processing prompts for multiple results (hiresfix)"
log( log(
self.ppp_logger, self.ppp_logger,
self.ppp_debug_level, self.ppp_debug_level,
logging.INFO, logging.INFO,
"processing prompts combinatorially (hiresfix)", msg,
) )
comb_results_hr = ppp.process_prompt( comb_results_hr = self.ppp.process_prompt(
rph[0], rph[0],
rnh[0], rnh[0],
seed_for_comb, calculated_seeds,
jobinfo={ jobinfo={
"job_timestamp": shared.state.job_timestamp, "job_timestamp": shared.state.job_timestamp,
"job": shared.state.job, "job": shared.state.job,
@@ -462,64 +532,86 @@ class PromptPostProcessorA1111Script(scripts.Script):
num_comb_hr = len(comb_results_hr) num_comb_hr = len(comb_results_hr)
for i in range(len(rph)): # pylint: disable=consider-using-enumerate for i in range(len(rph)): # pylint: disable=consider-using-enumerate
posp, negp, _ = comb_results_hr[i % num_comb_hr] 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))] if input_run_mode == RUN_MODE.combinatorial.value:
field_name = "PPP HR combination"
else:
field_name = "PPP HR result"
extra_params[field_name] = [str(1 + (i % num_comb_hr)) for i in range(len(rph))]
else: else:
# processes prompts # processes prompts
for prompttype, typeindex in prompts_list.keys(): for index, grouplist in enumerate(prompts_list):
log( for typeindex in range(len(grouplist)):
self.ppp_logger, typeprompt = [regular_type, hiresfix_type][typeindex]
self.ppp_debug_level, log(
logging.INFO, self.ppp_logger,
f"processing prompts ({prompttype}[{typeindex+1}])", self.ppp_debug_level,
) logging.INFO,
key = ( f"processing prompts ({typeprompt}[{index+1}])",
(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",
},
) )
posp, negp, _ = results[0] key = (
cached = (posp, negp) (hash_fullenv, calculated_seeds[index], rpr[index], rnr[index])
self.lru_cache.put(key, cached) if typeindex == 0
# adds also the result so i2i doesn't process it unnecessarily else (hash_fullenv, calculated_seeds[index], rph[index], rnh[index])
self.lru_cache.put((hsh, seed, posp, negp), cached) )
else: cached = self.lru_cache.get(key)
log(self.ppp_logger, self.ppp_debug_level, logging.INFO, "result already in cache") if cached is None:
prompts_list[(prompttype, typeindex)] = cached hsh, seed, prompt, negative_prompt = key
ppp.process_prompts_group_end() 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 = self.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
self.ppp.process_prompts_group_end()
# updates the prompts # updates the prompts
regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None) regular_copy = (rpr.copy() if rpr else None, rnr.copy() if rnr else None)
hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None) hiresfix_copy = (rph.copy() if rph else None, rnh.copy() if rnh else None)
regular_changes = False regular_changes = False
hiresfix_changes = False hiresfix_changes = False
for (prompttype, typeindex), (posp, negp) in prompts_list.items(): for index, grouplist in enumerate(prompts_list):
if prompttype == regular_type: for typeindex, groupprompts in enumerate(grouplist):
if rpr[typeindex].strip() != posp.strip() or rnr[typeindex].strip() != negp.strip(): if groupprompts is None:
regular_changes = True continue
rpr[typeindex] = posp posp, negp = groupprompts
rnr[typeindex] = negp if typeindex == 0:
elif prompttype == hiresfix_type: if rpr[index].strip() != posp.strip() or rnr[index].strip() != negp.strip():
if rph[typeindex].strip() != posp.strip() or rnh[typeindex].strip() != negp.strip(): regular_changes = True
hiresfix_changes = True rpr[index] = posp
rph[typeindex] = posp rnr[index] = negp
rnh[typeindex] = 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 # initialize extra generation parameters
if add_prompts: if add_prompts:
if hiresfix_exists:
extra_params["PPP Hires prompt"] = rph
extra_params["PPP Hires negative prompt"] = rnh
if regular_changes: if regular_changes:
extra_params["PPP original prompts"] = regular_copy[0] extra_params["PPP original prompts"] = regular_copy[0]
extra_params["PPP original negative prompts"] = regular_copy[1] extra_params["PPP original negative prompts"] = regular_copy[1]
+40 -30
View File
@@ -1,4 +1,4 @@
from dataclasses import replace from dataclasses import replace, make_dataclass
import difflib import difflib
import logging import logging
from pathlib import Path from pathlib import Path
@@ -6,7 +6,16 @@ from typing import Any, NamedTuple, Optional
import unittest import unittest
import datetime import datetime
from ppp_classes import IFWILDCARDS_CHOICES, ONWARNING_CHOICES, PPPStateOptions from ppp_classes import (
DEFAULT_SAMPLER,
IFWILDCARDS_CHOICES,
NEXT_SEED,
ONWARNING_CHOICES,
PPPEnvInfo,
RUN_MODE,
SUPPORTED_APPS,
PPPStateOptions,
)
from ppp_enmappings import PPPExtraNetworkMappings # type: ignore from ppp_enmappings import PPPExtraNetworkMappings # type: ignore
from ppp_wildcards import PPPWildcards # type: ignore from ppp_wildcards import PPPWildcards # type: ignore
from ppp import PromptPostProcessor # type: ignore from ppp import PromptPostProcessor # type: ignore
@@ -74,19 +83,22 @@ class TestPromptPostProcessorBase(unittest.TestCase):
cup_extranetwork_tags=True, cup_extranetwork_tags=True,
cup_merge_attention=True, cup_merge_attention=True,
cup_remove_extranetwork_tags=False, cup_remove_extranetwork_tags=False,
do_combinatorial=False,
combinatorial_shuffle=False,
combinatorial_limit=0,
results_file=(Path(__file__).parent / "logs" / "output_%date%.txt") if enable_file_logging else "", results_file=(Path(__file__).parent / "logs" / "output_%date%.txt") if enable_file_logging else "",
run_mode=RUN_MODE.single,
results_limit=0,
results_shuffle=False,
comb_random_fixed=True,
default_sampler=DEFAULT_SAMPLER.random,
next_seed=NEXT_SEED.randomize,
)
self.def_env_info = PPPEnvInfo(
app=SUPPORTED_APPS.tests,
ppp_config=None,
model_class="SDXL",
property_base=make_dataclass("PropertyBase", [("is_sdxl", bool)])(is_sdxl=True),
models_path="./webui/models",
model_filename="./webui/models/Stable-diffusion/testmodel.safetensors",
) )
self.def_env_info = {
"app": "tests",
"ppp_config": None,
"model_class": "SDXL",
"property_base": {"is_sdxl": True},
"models_path": "./webui/models",
"model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors",
}
self.interrupted = False self.interrupted = False
self.wildcards_obj = PPPWildcards(self.lf.log) self.wildcards_obj = PPPWildcards(self.lf.log)
self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log) self.extranetwork_maps_obj = PPPExtraNetworkMappings(self.lf.log)
@@ -119,8 +131,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
def init_ppp( def init_ppp(
self, self,
ppp: Optional[str | PromptPostProcessor] = None, ppp: Optional[str | PromptPostProcessor] = None,
combinatorial: bool = False, **kwargs,
combinatorial_limit: int = 0,
) -> PromptPostProcessor: ) -> PromptPostProcessor:
if isinstance(ppp, str): if isinstance(ppp, str):
if ppp == "nocup": if ppp == "nocup":
@@ -142,8 +153,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
cup_ands_eol=False, cup_ands_eol=False,
cup_extranetwork_tags=False, cup_extranetwork_tags=False,
cup_merge_attention=False, cup_merge_attention=False,
do_combinatorial=combinatorial, **kwargs,
combinatorial_limit=combinatorial_limit,
), ),
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -157,8 +167,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
replace( replace(
self.defopts, self.defopts,
strict_operators=False, strict_operators=False,
do_combinatorial=combinatorial, **kwargs,
combinatorial_limit=combinatorial_limit,
), ),
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -173,8 +182,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
self.def_env_info, self.def_env_info,
replace( replace(
self.defopts, self.defopts,
do_combinatorial=combinatorial, **kwargs,
combinatorial_limit=combinatorial_limit,
), ),
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -200,10 +208,9 @@ class TestPromptPostProcessorBase(unittest.TestCase):
seed: int = 1, seed: int = 1,
ppp: Optional[str | PromptPostProcessor] = None, ppp: Optional[str | PromptPostProcessor] = None,
interrupted: bool = False, interrupted: bool = False,
combinatorial: bool = False,
combinatorial_limit: int = 0,
specific_wc_folders: Optional[list[Path]] = None, specific_wc_folders: Optional[list[Path]] = None,
specific_em_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. Process the prompt and compare the results with the expected prompts.
@@ -214,10 +221,9 @@ class TestPromptPostProcessorBase(unittest.TestCase):
seed (int, optional): The seed value. Defaults to 1. seed (int, optional): The seed value. Defaults to 1.
ppp (Optional[str | PromptPostProcessor], optional): The PromptPostProcessor instance or type. Defaults to None. ppp (Optional[str | PromptPostProcessor], optional): The PromptPostProcessor instance or type. Defaults to None.
interrupted (bool, optional): The interrupted flag. Defaults to False. interrupted (bool, optional): The interrupted flag. Defaults to False.
combinatorial (bool, optional): The combinatorial flag. Defaults to False.
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_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. 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: Returns:
None None
@@ -232,13 +238,13 @@ class TestPromptPostProcessorBase(unittest.TestCase):
DEBUG_LEVEL.full, DEBUG_LEVEL.full,
specific_em_folders, specific_em_folders,
) )
the_obj: PromptPostProcessor = self.init_ppp(ppp, combinatorial, combinatorial_limit) the_obj: PromptPostProcessor = ppp if isinstance(ppp, PromptPostProcessor) else self.init_ppp(ppp)
out = ( out = (
[OutputTuple("", "", None)] [OutputTuple("", "", None)]
if expected_output is None if expected_output is None
else expected_output if isinstance(expected_output, list) else [expected_output] else expected_output if isinstance(expected_output, list) else [expected_output]
) )
if the_obj.state.options.do_combinatorial: if the_obj.state.options.run_mode in (RUN_MODE.multiple, RUN_MODE.combinatorial):
# combinatorial # combinatorial
errors = [] errors = []
the_obj.process_prompts_group_start() the_obj.process_prompts_group_start()
@@ -246,6 +252,8 @@ class TestPromptPostProcessorBase(unittest.TestCase):
input_prompts.prompt, input_prompts.prompt,
input_prompts.negative_prompt, input_prompts.negative_prompt,
seed, seed,
jobinfo={"test_case": self.id()},
input_vars=input_vars,
) )
the_obj.process_prompts_group_end() the_obj.process_prompts_group_end()
self.assertTrue( self.assertTrue(
@@ -254,7 +262,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
) )
if not self.interrupted and expected_output is not None: if not self.interrupted and expected_output is not None:
if len(result) != len(out): if len(result) != len(out):
errors.append(f"Incorrect number of combinations: got {len(result)} but expected {len(out)}") errors.append(f"Incorrect number of results: got {len(result)} but expected {len(out)}")
for out_prompt, out_negative_prompt, out_variables in out: for out_prompt, out_negative_prompt, out_variables in out:
found = None found = None
for r_prompt, r_negative_prompt, r_variables in result: for r_prompt, r_negative_prompt, r_variables in result:
@@ -264,7 +272,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
if not found: if not found:
errors.extend( errors.extend(
[ [
"Combination not found in output", "Result not found in output",
"Prompt:", "Prompt:",
out_prompt, out_prompt,
"Negative Prompt:", "Negative Prompt:",
@@ -286,7 +294,7 @@ class TestPromptPostProcessorBase(unittest.TestCase):
if missing_vars or incorrect_vars: if missing_vars or incorrect_vars:
errors.extend( errors.extend(
[ [
"Combination found, but variables do not match", "Result found, but variables do not match",
"Prompt:", "Prompt:",
out_prompt, out_prompt,
"Negative Prompt:", "Negative Prompt:",
@@ -311,6 +319,8 @@ class TestPromptPostProcessorBase(unittest.TestCase):
input_prompts.prompt, input_prompts.prompt,
input_prompts.negative_prompt, input_prompts.negative_prompt,
seed, seed,
jobinfo={"test_case": self.id()},
input_vars=input_vars,
) )
self.assertTrue( self.assertTrue(
self.interrupted == interrupted, self.interrupted == interrupted,
+135 -29
View File
@@ -1,6 +1,4 @@
from dataclasses import replace from ppp_classes import DEFAULT_SAMPLER, NEXT_SEED, RUN_MODE # type: ignore
from ppp import PromptPostProcessor # type: ignore
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__": if __name__ == "__main__":
@@ -22,7 +20,6 @@ class TestChoices(TestPromptPostProcessorBase):
) )
def test_ch_cyclical(self): # cyclical sampler cycles through all choices def test_ch_cyclical(self): # cyclical sampler cycles through all choices
ppp_instance = self.init_ppp("nocup")
self.process( self.process(
InputTuple("the choices are: {@choice1|choice2|choice3}", ""), InputTuple("the choices are: {@choice1|choice2|choice3}", ""),
[ [
@@ -31,11 +28,10 @@ class TestChoices(TestPromptPostProcessorBase):
OutputTuple("the choices are: choice3", ""), OutputTuple("the choices are: choice3", ""),
OutputTuple("the choices are: choice1", ""), # cycles back OutputTuple("the choices are: choice1", ""), # cycles back
], ],
ppp=ppp_instance, ppp="nocup",
) )
def test_ch_cyclical_multiple_constructs(self): # two independent @ constructs cycle together def test_ch_cyclical_multiple_constructs(self): # two independent @ constructs cycle together
ppp_instance = self.init_ppp("nocup")
self.process( self.process(
InputTuple("{@a|b} {@c|d}", ""), InputTuple("{@a|b} {@c|d}", ""),
[ [
@@ -45,7 +41,7 @@ class TestChoices(TestPromptPostProcessorBase):
OutputTuple("b d", ""), OutputTuple("b d", ""),
OutputTuple("a c", ""), # cycles back OutputTuple("a c", ""), # cycles back
], ],
ppp=ppp_instance, ppp="nocup",
) )
def test_ch_cyclical_resets_on_prompt_change(self): # state resets when the prompt pair changes def test_ch_cyclical_resets_on_prompt_change(self): # state resets when the prompt pair changes
@@ -67,7 +63,6 @@ class TestChoices(TestPromptPostProcessorBase):
) )
def test_ch_cyclical_mixed_samplers(self): # @ construct cycles while a ~ construct alongside is unaffected def test_ch_cyclical_mixed_samplers(self): # @ construct cycles while a ~ construct alongside is unaffected
ppp_instance = self.init_ppp("nocup")
self.process( self.process(
InputTuple("{@a|b|c} {x|y}", ""), InputTuple("{@a|b|c} {x|y}", ""),
[ [
@@ -76,7 +71,7 @@ class TestChoices(TestPromptPostProcessorBase):
OutputTuple("c x", ""), OutputTuple("c x", ""),
OutputTuple("a y", ""), # @ cycles back OutputTuple("a y", ""), # @ cycles back
], ],
ppp=ppp_instance, ppp="nocup",
) )
def test_ch_choices_withcomments(self): # choices with comments and multiline def test_ch_choices_withcomments(self): # choices with comments and multiline
@@ -103,6 +98,13 @@ class TestChoices(TestPromptPostProcessorBase):
ppp="nocup", ppp="nocup",
) )
def test_ch_choices_if_default(self): # choices with if and a default
self.process(
InputTuple("the choice is: {if false::choice1|if _is_sd1::choice2|else::choice3}", ""),
OutputTuple("the choice is: choice3", ""),
ppp="nocup",
)
def test_ch_choices_set_if_multiple(self): # choices with if user variable and multiple selection def test_ch_choices_set_if_multiple(self): # choices with if user variable and multiple selection
self.process( self.process(
InputTuple("${var=test}the choices are: {2$$, $$3::choice1|2 if not var eq 'test'::choice2|choice3}", ""), InputTuple("${var=test}the choices are: {2$$, $$3::choice1|2 if not var eq 'test'::choice2|choice3}", ""),
@@ -131,18 +133,7 @@ class TestChoices(TestPromptPostProcessorBase):
self.process( self.process(
InputTuple("<lora:test1:1><lora:test2:{0.2|0.5|0.7|1}>", ""), InputTuple("<lora:test1:1><lora:test2:{0.2|0.5|0.7|1}>", ""),
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=self.init_ppp(None, cup_remove_extranetwork_tags=True),
self.ppp_logger,
self.def_env_info,
replace(
self.defopts,
cup_remove_extranetwork_tags=True,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
),
) )
def test_ch_cmd_includewildcard(self): def test_ch_cmd_includewildcard(self):
@@ -156,14 +147,129 @@ class TestChoices(TestPromptPostProcessorBase):
def test_ch_combinatorial(self): def test_ch_combinatorial(self):
self.process( self.process(
InputTuple("{choice1|choice2|choice3}, ${v:{option1|option2}}", ""), InputTuple("{choice1|choice2|choice3}, ${v:{option1|option2}}, {~a|b}", ""),
[ [
OutputTuple("choice1, option1", ""), OutputTuple("choice1, option1, a", ""),
OutputTuple("choice1, option2", ""), OutputTuple("choice1, option2, b", ""),
OutputTuple("choice2, option1", ""), OutputTuple("choice2, option1, a", ""),
OutputTuple("choice2, option2", ""), OutputTuple("choice2, option2, b", ""),
OutputTuple("choice3, option1", ""), OutputTuple("choice3, option1, b", ""),
OutputTuple("choice3, option2", "", {"v": "option2"}), OutputTuple("choice3, option2, a", "", {"v": "option2"}),
], ],
combinatorial=True, ppp=self.init_ppp(
None,
run_mode=RUN_MODE.combinatorial,
comb_random_fixed=False, # allow different random choices across combinations
),
) )
def test_ch_comb_random_consistent(self): # ~ sampler picks one value shared across all combinations
ppp_instance = self.init_ppp("nocup", run_mode=RUN_MODE.combinatorial)
ppp_instance.process_prompts_group_start()
result = ppp_instance.process_prompt("{~a|b|c} {x|y}", "", starting_seed=1)
ppp_instance.process_prompts_group_end()
self.assertEqual(len(result), 2, "Expected exactly 2 combinations ({x|y} expands to 2)")
rnd_choices = {r_prompt.split()[0] for r_prompt, _, _ in result}
self.assertEqual(
len(rnd_choices),
1,
f"The ~ sampler must yield the same value across all combinations, got: {rnd_choices}",
)
# Multiple
def test_ch_multiple(self):
self.process(
InputTuple("{choice1|choice2|choice3}, ${v:{option1|option2}}, {~a|b}", ""),
[
OutputTuple("choice2, option2, a", ""),
OutputTuple("choice1, option1, b", ""),
OutputTuple("choice1, option2, b", ""),
OutputTuple("choice3, option1, b", ""),
OutputTuple("choice1, option1, a", ""),
OutputTuple("choice2, option1, a", "", {"v": "option1"}),
],
ppp=self.init_ppp(
None,
run_mode=RUN_MODE.multiple,
results_limit=6,
),
)
# Default sampler
def test_ch_default_sampler_cyclical_single(self):
self.process(
InputTuple("{choice1|choice2|choice3}", ""),
[
OutputTuple("choice1", ""),
OutputTuple("choice2", ""),
OutputTuple("choice3", ""),
OutputTuple("choice1", ""),
],
ppp=self.init_ppp("nocup", default_sampler=DEFAULT_SAMPLER.cyclical),
)
def test_ch_default_sampler_cyclical_multiple(self):
self.process(
InputTuple("{choice1|choice2|choice3}", ""),
[
OutputTuple("choice1", ""),
OutputTuple("choice2", ""),
OutputTuple("choice3", ""),
OutputTuple("choice1", ""),
],
ppp=self.init_ppp(
"nocup",
default_sampler=DEFAULT_SAMPLER.cyclical,
run_mode=RUN_MODE.multiple,
results_limit=4,
),
)
# next_seed / _output_seed
def test_ch_next_seed_input(self): # input mode keeps the same seed for every result
self.process(
InputTuple("{@a|b|c}", ""),
[
OutputTuple("a", "", {"_output_seed": 1}),
OutputTuple("b", "", {"_output_seed": 1}),
OutputTuple("c", "", {"_output_seed": 1}),
],
seed=1,
ppp=self.init_ppp("nocup", run_mode=RUN_MODE.multiple, results_limit=3, next_seed=NEXT_SEED.input),
)
def test_ch_next_seed_increment(self): # increment mode increases the seed by 1 for each result
self.process(
InputTuple("{@a|b|c}", ""),
[
OutputTuple("a", "", {"_output_seed": 1}),
OutputTuple("b", "", {"_output_seed": 2}),
OutputTuple("c", "", {"_output_seed": 3}),
],
seed=1,
ppp=self.init_ppp("nocup", run_mode=RUN_MODE.multiple, results_limit=3, next_seed=NEXT_SEED.increment),
)
def test_ch_next_seed_decrement(self): # decrement mode decreases the seed by 1 for each result
self.process(
InputTuple("{@a|b|c}", ""),
[
OutputTuple("a", "", {"_output_seed": 3}),
OutputTuple("b", "", {"_output_seed": 2}),
OutputTuple("c", "", {"_output_seed": 1}),
],
seed=3,
ppp=self.init_ppp("nocup", run_mode=RUN_MODE.multiple, results_limit=3, next_seed=NEXT_SEED.decrement),
)
def test_ch_next_seed_randomize(self): # randomize mode produces a distinct seed for every result
ppp_instance = self.init_ppp("nocup", run_mode=RUN_MODE.multiple, results_limit=3, next_seed=NEXT_SEED.randomize)
ppp_instance.process_prompts_group_start()
results = ppp_instance.process_prompt("{@a|b|c}", "", starting_seed=1)
ppp_instance.process_prompts_group_end()
seeds = [r_vars.get("_output_seed") for _, _, r_vars in results]
self.assertEqual(len(seeds), 3, f"Expected 3 results, got {len(seeds)}")
self.assertEqual(len(set(seeds)), 3, f"Expected 3 distinct seeds, got: {seeds}")
+21 -49
View File
@@ -1,10 +1,7 @@
import logging import logging
from dataclasses import replace
from ppp import PromptPostProcessor # type: ignore
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__": if __name__ == "__main__":
raise SystemExit("This script must not be run directly") raise SystemExit("This script must not be run directly")
@@ -38,36 +35,17 @@ class TestCleanup(TestPromptPostProcessorBase):
self.process( self.process(
InputTuple("this is a <lora:test:1> test__yaml/wildcard7__", ""), InputTuple("this is a <lora:test:1> test__yaml/wildcard7__", ""),
OutputTuple("this is a test", ""), OutputTuple("this is a test", ""),
ppp=PromptPostProcessor( ppp=self.init_ppp(None, cup_remove_extranetwork_tags=True),
self.ppp_logger,
self.def_env_info,
replace(
self.defopts,
cup_remove_extranetwork_tags=True,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
),
) )
def test_cl_dontremoveseparatorsoneol(self): # don't remove separators on eol def test_cl_dontremoveseparatorsoneol(self): # don't remove separators on eol
self.process( self.process(
InputTuple("this is a test,\nsecond line", ""), InputTuple("this is a test,\nsecond line", ""),
OutputTuple("this is a test,\nsecond line", ""), OutputTuple("this is a test,\nsecond line", ""),
ppp=PromptPostProcessor( ppp=self.init_ppp(
self.ppp_logger, None,
self.def_env_info, cup_extra_separators2=False,
replace( cup_extra_separators_include_eol=False,
self.defopts,
cup_extra_separators2=False,
cup_extra_separators_include_eol=False,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
), ),
) )
@@ -85,27 +63,19 @@ class TestCleanup(TestPromptPostProcessorBase):
(d:0.9)""", (d:0.9)""",
"", "",
), ),
ppp=PromptPostProcessor( ppp=self.init_ppp(
self.ppp_logger, None,
self.def_env_info, cup_empty_constructs=False,
replace( cup_extra_separators=True,
self.defopts, cup_extra_separators2=False,
cup_empty_constructs=False, cup_extra_separators_include_eol=False,
cup_extra_separators=True, cup_extra_spaces=False,
cup_extra_separators2=False, cup_breaks=False,
cup_extra_separators_include_eol=False, cup_breaks_eol=False,
cup_extra_spaces=False, cup_ands=False,
cup_breaks=False, cup_ands_eol=False,
cup_breaks_eol=False, cup_extranetwork_tags=False,
cup_ands=False, cup_merge_attention=False,
cup_ands_eol=False,
cup_extranetwork_tags=False,
cup_merge_attention=False,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
), ),
) )
@@ -191,7 +161,9 @@ class TestCleanup(TestPromptPostProcessorBase):
"Expected an 'Unmatched' warning", "Expected an 'Unmatched' warning",
) )
def test_cl_warn_escaped_unmatched_no_false_warning(self): # escaped unmatched paren/bracket does not trigger warning def test_cl_warn_escaped_unmatched_no_false_warning(
self,
): # escaped unmatched paren/bracket does not trigger warning
with self.assertNoLogs("PromptPostProcessor", level=logging.WARNING): with self.assertNoLogs("PromptPostProcessor", level=logging.WARNING):
self.process( self.process(
InputTuple(r"text with \(escaped unmatched\]", ""), InputTuple(r"text with \(escaped unmatched\]", ""),
+81 -80
View File
@@ -1,3 +1,4 @@
from dataclasses import replace
from ppp import PromptPostProcessor # type: ignore from ppp import PromptPostProcessor # type: ignore
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
@@ -21,10 +22,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""), OutputTuple("(test1:0.9) (test2) (test3:1.5) (test4:0.99)", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"attention": "parentheses"}}}, ppp_config={"hosts": {"tests": {"attention": "parentheses"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -42,10 +43,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1 test2 test3", ""), OutputTuple("test1 test2 test3", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"attention": "disable"}}}, ppp_config={"hosts": {"tests": {"attention": "disable"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -63,10 +64,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"attention": "remove"}}}, ppp_config={"hosts": {"tests": {"attention": "remove"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -84,10 +85,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"attention": "error"}}}, ppp_config={"hosts": {"tests": {"attention": "error"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -106,10 +107,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1", ""), OutputTuple("test1", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"scheduling": "before"}}}, ppp_config={"hosts": {"tests": {"scheduling": "before"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -127,10 +128,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test2", ""), OutputTuple("test2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"scheduling": "after"}}}, ppp_config={"hosts": {"tests": {"scheduling": "after"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -148,10 +149,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1 test3", ""), OutputTuple("test1 test3", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"scheduling": "first"}}}, ppp_config={"hosts": {"tests": {"scheduling": "first"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -169,10 +170,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"scheduling": "remove"}}}, ppp_config={"hosts": {"tests": {"scheduling": "remove"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -190,10 +191,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"scheduling": "error"}}}, ppp_config={"hosts": {"tests": {"scheduling": "error"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -212,10 +213,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1", ""), OutputTuple("test1", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"alternation": "first"}}}, ppp_config={"hosts": {"tests": {"alternation": "first"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -233,10 +234,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"alternation": "remove"}}}, ppp_config={"hosts": {"tests": {"alternation": "remove"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -254,10 +255,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"alternation": "error"}}}, ppp_config={"hosts": {"tests": {"alternation": "error"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -276,10 +277,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1\ntest2", ""), OutputTuple("test1\ntest2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"and": "eol"}}}, ppp_config={"hosts": {"tests": {"and": "eol"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -297,10 +298,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1, test2", ""), OutputTuple("test1, test2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"and": "comma"}}}, ppp_config={"hosts": {"tests": {"and": "comma"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -318,10 +319,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1 test2", ""), OutputTuple("test1 test2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"and": "remove"}}}, ppp_config={"hosts": {"tests": {"and": "remove"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -339,10 +340,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"and": "error"}}}, ppp_config={"hosts": {"tests": {"and": "error"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -361,10 +362,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1\ntest2", ""), OutputTuple("test1\ntest2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"break": "eol"}}}, ppp_config={"hosts": {"tests": {"break": "eol"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -382,10 +383,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1, test2", ""), OutputTuple("test1, test2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"break": "comma"}}}, ppp_config={"hosts": {"tests": {"break": "comma"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -403,10 +404,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("test1 test2", ""), OutputTuple("test1 test2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"break": "remove"}}}, ppp_config={"hosts": {"tests": {"break": "remove"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -424,10 +425,10 @@ class TestHosts(TestPromptPostProcessorBase):
OutputTuple("", ""), OutputTuple("", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"ppp_config": {"hosts": {"tests": {"break": "error"}}}, ppp_config={"hosts": {"tests": {"break": "error"}}},
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
+38 -68
View File
@@ -1,5 +1,4 @@
from dataclasses import replace from dataclasses import replace
from ppp import PromptPostProcessor # type: ignore from ppp import PromptPostProcessor # type: ignore
from ppp_classes import ONWARNING_CHOICES # type: ignore from ppp_classes import ONWARNING_CHOICES # type: ignore
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
@@ -57,18 +56,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"", "",
), ),
OutputTuple("", "", {"v1": ""}), OutputTuple("", "", {"v1": ""}),
ppp=PromptPostProcessor( ppp=self.init_ppp(None, on_warning=ONWARNING_CHOICES.warn),
self.ppp_logger,
self.def_env_info,
replace(
self.defopts,
on_warning=ONWARNING_CHOICES.warn,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
),
) )
# Variable in extranetworks # Variable in extranetworks
@@ -690,18 +678,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"", "",
), ),
OutputTuple("not OK", ""), OutputTuple("not OK", ""),
ppp=PromptPostProcessor( ppp=self.init_ppp(None, on_warning=ONWARNING_CHOICES.warn),
self.ppp_logger,
self.def_env_info,
replace(
self.defopts,
on_warning=ONWARNING_CHOICES.warn,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
),
) )
def test_cmd_if_undefined_var_int_compare_stop(self): # undefined var integer compare with on_warning=stop def test_cmd_if_undefined_var_int_compare_stop(self): # undefined var integer compare with on_warning=stop
@@ -721,18 +698,7 @@ class TestVarCommands(TestPromptPostProcessorBase):
"", "",
), ),
OutputTuple("not OK", ""), OutputTuple("not OK", ""),
ppp=PromptPostProcessor( ppp=self.init_ppp(None, on_warning=ONWARNING_CHOICES.warn),
self.ppp_logger,
self.def_env_info,
replace(
self.defopts,
on_warning=ONWARNING_CHOICES.warn,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
),
) )
def test_cmd_if_nonnumeric_var_int_compare_stop(self): # non-numeric var integer compare with on_warning=stop def test_cmd_if_nonnumeric_var_int_compare_stop(self): # non-numeric var integer compare with on_warning=stop
@@ -752,18 +718,22 @@ class TestVarCommands(TestPromptPostProcessorBase):
"", "",
), ),
OutputTuple("not OK", ""), OutputTuple("not OK", ""),
ppp=PromptPostProcessor( ppp=self.init_ppp(None, on_warning=ONWARNING_CHOICES.warn),
self.ppp_logger, )
self.def_env_info,
replace( # Input variables
self.defopts,
on_warning=ONWARNING_CHOICES.warn, def test_input_variables(self):
), self.process(
self.grammar_content, InputTuple(
self.interrupt, "${_input_prev_positive_prompt}, high quality",
self.wildcards_obj, "",
self.extranetwork_maps_obj,
), ),
OutputTuple(
"this is a test, high quality",
"",
),
input_vars={"prev_positive_prompt": "this is a test"},
) )
# Command tests # Command tests
@@ -801,10 +771,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
OutputTuple("this is PONY", ""), OutputTuple("this is PONY", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -1034,10 +1004,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
OutputTuple("<lora:lorapony:0.8>inlinetrigger, triggerpony1, triggerpony2", ""), OutputTuple("<lora:lorapony:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -1055,10 +1025,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
OutputTuple("<lora:lorapony:0.4>inlinetrigger, triggerpony1, triggerpony2", ""), OutputTuple("<lora:lorapony:0.4>inlinetrigger, triggerpony1, triggerpony2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -1076,10 +1046,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
OutputTuple("<lora:lorapony:0.6:0.8>inlinetrigger, triggerpony1, triggerpony2", ""), OutputTuple("<lora:lorapony:0.6:0.8>inlinetrigger, triggerpony1, triggerpony2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
@@ -1097,10 +1067,10 @@ class TestVarCommands(TestPromptPostProcessorBase):
OutputTuple("<lora:loraillustrious:0.9:0.8>inlinetrigger, triggerillustrious1, triggerillustrious2", ""), OutputTuple("<lora:loraillustrious:0.9:0.8>inlinetrigger, triggerillustrious1, triggerillustrious2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"model_filename": "./webui/models/Stable-diffusion/ilxlmodel.safetensors", model_filename="./webui/models/Stable-diffusion/ilxlmodel.safetensors",
}, ),
self.defopts, self.defopts,
self.grammar_content, self.grammar_content,
self.interrupt, self.interrupt,
+10 -10
View File
@@ -24,10 +24,10 @@ class TestModelVariants(TestPromptPostProcessorBase):
OutputTuple("test1test2", ""), OutputTuple("test1test2", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"model_filename": "./webui/models/Stable-diffusion/testmodel.safetensors", model_filename="./webui/models/Stable-diffusion/testmodel.safetensors",
"ppp_config": { ppp_config={
"models": { "models": {
"sd1": { "sd1": {
"detect": {"tests": {"class": ["SD15", "SD15_instructpix2pix"]}}, "detect": {"tests": {"class": ["SD15", "SD15_instructpix2pix"]}},
@@ -62,7 +62,7 @@ class TestModelVariants(TestPromptPostProcessorBase):
}, },
} }
}, },
}, ),
replace( replace(
self.defopts, self.defopts,
on_warning=ONWARNING_CHOICES.warn, on_warning=ONWARNING_CHOICES.warn,
@@ -84,15 +84,15 @@ class TestModelVariants(TestPromptPostProcessorBase):
OutputTuple("not SDXL, not PONY", ""), OutputTuple("not SDXL, not PONY", ""),
ppp=PromptPostProcessor( ppp=PromptPostProcessor(
self.ppp_logger, self.ppp_logger,
{ replace(
**self.def_env_info, self.def_env_info,
"model_filename": "./webui/models/Stable-diffusion/ponymodel.safetensors", model_filename="./webui/models/Stable-diffusion/ponymodel.safetensors",
"ppp_config": { ppp_config={
"models": { "models": {
"sdxl": None, "sdxl": None,
} }
}, },
}, ),
replace( replace(
self.defopts, self.defopts,
on_warning=ONWARNING_CHOICES.warn, on_warning=ONWARNING_CHOICES.warn,
+54 -83
View File
@@ -1,7 +1,6 @@
from dataclasses import replace
from ppp import PromptPostProcessor from ppp import PromptPostProcessor
from ppp_classes import IFWILDCARDS_CHOICES from ppp_classes import IFWILDCARDS_CHOICES, RUN_MODE
from ppp_logging import DEBUG_LEVEL
from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase from .base_tests import OutputTuple, InputTuple, TestPromptPostProcessorBase
if __name__ == "__main__": if __name__ == "__main__":
@@ -19,18 +18,10 @@ class TestWildcards(TestPromptPostProcessorBase):
self.process( self.process(
InputTuple("__bad_wildcard__", "{option1|option2}"), InputTuple("__bad_wildcard__", "{option1|option2}"),
OutputTuple("__bad_wildcard__", "{option1|option2}"), OutputTuple("__bad_wildcard__", "{option1|option2}"),
ppp=PromptPostProcessor( ppp=self.init_ppp(
self.ppp_logger, None,
self.def_env_info, process_wildcards=False,
replace( if_wildcards=IFWILDCARDS_CHOICES.ignore,
self.defopts,
process_wildcards=False,
if_wildcards=IFWILDCARDS_CHOICES.ignore,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
), ),
) )
@@ -44,18 +35,10 @@ class TestWildcards(TestPromptPostProcessorBase):
"this is: a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]<lora:xxx:1>", "this is: a (([complex|simple|regular] test)(test:2):1.5)\nBREAK with [abc:def:5]<lora:xxx:1>",
"[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]", "[neg5], ([|neg6|]:1.65), (neg1:1.65), [neg4::5], normal quality, [neg2(neg3:1.6):5]",
), ),
ppp=PromptPostProcessor( ppp=self.init_ppp(
self.ppp_logger, None,
self.def_env_info, process_wildcards=False,
replace( if_wildcards=IFWILDCARDS_CHOICES.remove,
self.defopts,
process_wildcards=False,
if_wildcards=IFWILDCARDS_CHOICES.remove,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
), ),
) )
@@ -63,18 +46,10 @@ class TestWildcards(TestPromptPostProcessorBase):
self.process( self.process(
InputTuple("__bad_wildcard__", "{option1|option2}"), InputTuple("__bad_wildcard__", "{option1|option2}"),
OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"), OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", "{option1|option2}"),
ppp=PromptPostProcessor( ppp=self.init_ppp(
self.ppp_logger, None,
self.def_env_info, process_wildcards=False,
replace( if_wildcards=IFWILDCARDS_CHOICES.warn,
self.defopts,
process_wildcards=False,
if_wildcards=IFWILDCARDS_CHOICES.warn,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
), ),
) )
@@ -85,18 +60,10 @@ class TestWildcards(TestPromptPostProcessorBase):
PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__bad_wildcard__", PromptPostProcessor.WILDCARD_STOP.format("__bad_wildcard__") + "__bad_wildcard__",
"{option1|option2}", "{option1|option2}",
), ),
ppp=PromptPostProcessor( ppp=self.init_ppp(
self.ppp_logger, None,
self.def_env_info, process_wildcards=False,
replace( if_wildcards=IFWILDCARDS_CHOICES.stop,
self.defopts,
process_wildcards=False,
if_wildcards=IFWILDCARDS_CHOICES.stop,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
), ),
interrupted=True, interrupted=True,
) )
@@ -105,18 +72,10 @@ class TestWildcards(TestPromptPostProcessorBase):
self.process( self.process(
InputTuple("${v=__bad_wildcard__}${v}", ""), InputTuple("${v=__bad_wildcard__}${v}", ""),
OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""), OutputTuple(PromptPostProcessor.WILDCARD_WARNING + "__bad_wildcard__", ""),
ppp=PromptPostProcessor( ppp=self.init_ppp(
self.ppp_logger, None,
self.def_env_info, process_wildcards=False,
replace( if_wildcards=IFWILDCARDS_CHOICES.warn,
self.defopts,
process_wildcards=False,
if_wildcards=IFWILDCARDS_CHOICES.warn,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
), ),
) )
@@ -128,6 +87,22 @@ class TestWildcards(TestPromptPostProcessorBase):
interrupted=True, interrupted=True,
) )
def test_wc_load_invalid_wildcards(self):
self.wildcards_obj.refresh_wildcards(
DEBUG_LEVEL.full,
[],
"""
inv alid1:
- choice1
- choice2
inv.alid2:
- choice1
- choice2
""",
)
self.assertFalse(self.wildcards_obj.get_wildcards("inv alid1"), "invalid wildcard name should not be accepted")
self.assertFalse(self.wildcards_obj.get_wildcards("inv.alid2"), "invalid wildcard name should not be accepted")
def test_wc_wildcard1a_text(self): # simple text wildcard def test_wc_wildcard1a_text(self): # simple text wildcard
self.process( self.process(
InputTuple("the choices are: __text/wildcard1__", ""), InputTuple("the choices are: __text/wildcard1__", ""),
@@ -337,6 +312,13 @@ class TestWildcards(TestPromptPostProcessorBase):
ppp="nocup", ppp="nocup",
) )
def test_wc_wildcard_emptydefault(self): # empty wildcard with default
self.process(
InputTuple("the choices are: __yaml/empty_default_wildcard__", ""),
OutputTuple("the choices are: 6", ""),
ppp="nocup",
)
def test_wc_wildcard4_yaml(self): # simple yaml wildcard with one option def test_wc_wildcard4_yaml(self): # simple yaml wildcard with one option
self.process( self.process(
InputTuple("the choices are: __yaml/wildcard4__", ""), InputTuple("the choices are: __yaml/wildcard4__", ""),
@@ -492,7 +474,7 @@ class TestWildcards(TestPromptPostProcessorBase):
OutputTuple("the choices are: choice3, choice2, option1", "", {"v": "option1"}), OutputTuple("the choices are: choice3, choice2, option1", "", {"v": "option1"}),
OutputTuple("the choices are: choice3, choice2, option2", "", {"v": "option2"}), OutputTuple("the choices are: choice3, choice2, option2", "", {"v": "option2"}),
], ],
combinatorial=True, ppp=self.init_ppp(None, run_mode=RUN_MODE.combinatorial),
) )
def test_wc_combinatorial_2(self): # combinatorial wildcard def test_wc_combinatorial_2(self): # combinatorial wildcard
@@ -545,8 +527,7 @@ class TestWildcards(TestPromptPostProcessorBase):
OutputTuple("choice1-choice3", ""), OutputTuple("choice1-choice3", ""),
OutputTuple("choice3-choice1", ""), OutputTuple("choice3-choice1", ""),
], ],
ppp="nocup", ppp=self.init_ppp("nocup", run_mode=RUN_MODE.combinatorial),
combinatorial=True,
) )
def test_wc_combinatorial_3(self): # combinatorial wildcard (keep choice order) def test_wc_combinatorial_3(self): # combinatorial wildcard (keep choice order)
@@ -564,19 +545,11 @@ class TestWildcards(TestPromptPostProcessorBase):
## choices 1 and 3 ## choices 1 and 3
OutputTuple("choice1-choice3", ""), OutputTuple("choice1-choice3", ""),
], ],
ppp=PromptPostProcessor( ppp=self.init_ppp(
self.ppp_logger, None,
self.def_env_info, keep_choices_order=True,
replace( cup_do_cleanup=False,
self.defopts, run_mode=RUN_MODE.combinatorial,
keep_choices_order=True,
cup_do_cleanup=False,
do_combinatorial=True,
),
self.grammar_content,
self.interrupt,
self.wildcards_obj,
self.extranetwork_maps_obj,
), ),
) )
@@ -603,8 +576,7 @@ class TestWildcards(TestPromptPostProcessorBase):
OutputTuple("choice1-choice3", ""), OutputTuple("choice1-choice3", ""),
OutputTuple("choice3-choice1", ""), OutputTuple("choice3-choice1", ""),
], ],
ppp="nocup", ppp=self.init_ppp("nocup", run_mode=RUN_MODE.combinatorial),
combinatorial=True,
) )
def test_wc_combinatorial_5(self): # combinatorial nested wildcards and multiselection enmappings def test_wc_combinatorial_5(self): # combinatorial nested wildcards and multiselection enmappings
@@ -626,6 +598,5 @@ class TestWildcards(TestPromptPostProcessorBase):
OutputTuple("<lora:loraany1:0.8> trigger1, trigger2, ", ""), OutputTuple("<lora:loraany1:0.8> trigger1, trigger2, ", ""),
OutputTuple("<lora:loraany2:1> trigger3, trigger4, ", ""), OutputTuple("<lora:loraany2:1> trigger3, trigger4, ", ""),
], ],
ppp="nocup", ppp=self.init_ppp("nocup", run_mode=RUN_MODE.combinatorial),
combinatorial=True,
) )
+9 -1
View File
@@ -6,7 +6,7 @@ yaml:
- choice3 - choice3
wildcard2: wildcard2:
- ~r2-3'Wildcard description'$$-$$ - r2-3'Wildcard description'$$-$$
- "'label1,label2'4::choice1" - "'label1,label2'4::choice1"
- "3:: choice2 " - "3:: choice2 "
- { labels: ["label1", "label3"], weight: 2, content: choice3 } - { labels: ["label1", "label3"], weight: 2, content: choice3 }
@@ -111,6 +111,14 @@ yaml:
- if _sd in ("test1", "test2")::4 - if _sd in ("test1", "test2")::4
- if (false or false)::5 - if (false or false)::5
empty_default_wildcard:
- if false::1
- if false::2
- if false::3
- if _sd in ("test1", "test2")::4
- if (false or false)::5
- else::6
circular1: circular1:
- 5::__yaml/circular2__ - 5::__yaml/circular2__
- choice1 - choice1
+285
View File
@@ -0,0 +1,285 @@
#!/usr/bin/env python3
"""
Scans wildcard and extranetwork mapping files for LoRA references and checks
if those LoRAs exist in the specified folders.
Recognized LoRA reference formats:
<lora:NAME:weight> Standard A1111/ComfyUI inline LoRA tag.
<ppp:ext lora NAME ...> PPP explicit LoRA command (quoted or unquoted name).
LoRA names that contain wildcards or inline choices (e.g. __path__, {a|b}) are
reported as dynamic and skipped - they cannot be resolved statically.
Usage:
python check_loras.py -l LORA_FOLDER [LORA_FOLDER ...] [options]
Options:
-w, --wildcards One or more wildcard folder paths to scan.
-e, --enmappings One or more enmapping folder paths to scan.
-l, --loras One or more folders to search for LoRA files (required).
--extensions LoRA file extensions (default: .safetensors .pt .ckpt .bin).
--case-sensitive Enable case-sensitive name matching (default: case-insensitive).
-v, --verbose Also list LoRAs that were found.
"""
import argparse
import re
import sys
from collections import defaultdict
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
try:
from ruamel.yaml import YAML
except ImportError:
print("Error: ruamel.yaml is required. Install with: pip install ruamel.yaml", file=sys.stderr)
sys.exit(1)
WILDCARD_EXTENSIONS = {".yaml", ".yml", ".json", ".txt"}
ENMAPPING_EXTENSIONS = {".yaml", ".yml", ".json"}
# Matches <lora:NAME>, <lora:NAME:w>, <lora:NAME:w1:w2>
# The name ends at the first colon or closing >, but may contain spaces.
_RE_LORA_STANDARD = re.compile(r"<lora:([^:>\s][^:>]*?)(?::[^>]*)?>")
# Matches <ppp:ext lora NAME ...> where NAME is single-quoted, double-quoted, or unquoted.
# The unquoted form stops before whitespace, >, or / to avoid capturing the rest of the tag.
_RE_LORA_PPP_EXT = re.compile(r"<ppp:ext\s+lora\s+(?:'([^']*)'|\"([^\"]*)\"|([\w.\-()][^\s>/'\"]*))")
def _read_text(path: Path) -> str | None:
for encoding in ("utf-8", "cp1252"):
try:
return path.read_text(encoding=encoding)
except (UnicodeDecodeError, OSError):
continue
print(f"Warning: could not read file: {path}", file=sys.stderr)
return None
def _is_dynamic(name: str) -> bool:
# Names with choice syntax {a|b} or wildcard references __path__ cannot be
# statically resolved.
return "{" in name or "}" in name or name.count("__") >= 2
_RE_SD_ESCAPE = re.compile(r"\\(.)")
def _unescape_sd(name: str) -> str:
# SD prompts escape special characters with a backslash (e.g. \( \) \[ \] \: \\).
# The actual filename on disk has no escaping, so strip them before lookup.
return _RE_SD_ESCAPE.sub(r"\1", name)
def _extract_lora_names(text: str) -> list[str]:
names = []
for m in _RE_LORA_STANDARD.finditer(text):
names.append(_unescape_sd(m.group(1).strip()))
for m in _RE_LORA_PPP_EXT.finditer(text):
name = m.group(1) or m.group(2) or m.group(3)
if name:
names.append(_unescape_sd(name.strip()))
return names
def _walk_strings_with_path(value, path: str = "") -> list[tuple[str, str]]:
"""Recursively collect (string_value, key_path) pairs from a parsed YAML/JSON structure."""
if isinstance(value, str):
return [(value, path)]
if isinstance(value, list):
result = []
for item in value:
result.extend(_walk_strings_with_path(item, path))
return result
if isinstance(value, dict):
result = []
for k, v in value.items():
child_path = f"{path} > {k}" if path else str(k)
result.extend(_walk_strings_with_path(v, child_path))
return result
return []
def scan_wildcard_file(path: Path) -> list[tuple[str, str]]:
"""Return (lora_name, location) pairs found in a wildcard file."""
text = _read_text(path)
if text is None:
return []
found = []
suffix = path.suffix.lower()
if suffix == ".txt":
for lineno, line in enumerate(text.splitlines(), start=1):
stripped = line.strip()
if not stripped or stripped.startswith("#"):
continue
for name in _extract_lora_names(stripped):
found.append((name, f"line {lineno}"))
return found
yaml = YAML(typ="safe")
try:
data = yaml.load(text)
except Exception as exc: # pylint: disable=broad-except
print(f"Warning: could not parse {path}: {exc}", file=sys.stderr)
return []
for string_val, key_path in _walk_strings_with_path(data):
for name in _extract_lora_names(string_val):
found.append((name, key_path))
return found
def scan_enmapping_file(path: Path) -> list[tuple[str, str]]:
"""Return (lora_name, context_snippet) pairs found in an enmapping file."""
text = _read_text(path)
if text is None:
return []
yaml = YAML(typ="safe")
try:
data = yaml.load(text)
except Exception as exc: # pylint: disable=broad-except
print(f"Warning: could not parse {path}: {exc}", file=sys.stderr)
return []
if not isinstance(data, dict):
return []
found = []
lora_section = data.get("lora", {})
if not isinstance(lora_section, dict):
return []
for mapping_key, variants in lora_section.items():
if not isinstance(variants, list):
continue
for variant in variants:
if not isinstance(variant, dict):
continue
name = variant.get("name")
if name and isinstance(name, str):
found.append((name.strip(), f"lora > {mapping_key}"))
return found
def build_lora_index(lora_folders: list[Path], extensions: set[str], case_sensitive: bool) -> set[str]:
index: set[str] = set()
for folder in lora_folders:
if not folder.is_dir():
print(f"Warning: LoRA folder does not exist or is not a directory: {folder}", file=sys.stderr)
continue
for f in folder.rglob("*"):
if f.is_file() and f.suffix.lower() in extensions:
stem = f.stem if case_sensitive else f.stem.lower()
index.add(stem)
return index
def main() -> None:
parser = argparse.ArgumentParser(
description="Check that LoRAs referenced in wildcard and enmapping files exist on disk.",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog=__doc__,
)
parser.add_argument(
"-w", "--wildcards",
nargs="*", type=Path, default=[], metavar="FOLDER",
help="Wildcard folder(s) to scan.",
)
parser.add_argument(
"-e", "--enmappings",
nargs="*", type=Path, default=[], metavar="FOLDER",
help="Enmapping folder(s) to scan.",
)
parser.add_argument(
"-l", "--loras",
nargs="+", type=Path, required=True, metavar="FOLDER",
help="Folder(s) to search for LoRA files.",
)
parser.add_argument(
"--extensions",
nargs="+", default=[".safetensors", ".pt", ".ckpt", ".bin"], metavar="EXT",
help="LoRA file extensions to recognize (default: .safetensors .pt .ckpt .bin).",
)
parser.add_argument(
"--case-sensitive",
action="store_true",
help="Enable case-sensitive name matching (default: case-insensitive).",
)
parser.add_argument(
"-v", "--verbose",
action="store_true",
help="Also list LoRAs that were found.",
)
args = parser.parse_args()
if not args.wildcards and not args.enmappings:
print("Error: specify at least one --wildcards or --enmappings folder.", file=sys.stderr)
sys.exit(1)
extensions = {(ext if ext.startswith(".") else f".{ext}").lower() for ext in args.extensions}
lora_index = build_lora_index(args.loras, extensions, args.case_sensitive)
# Keyed by lora name; value is a list of (source_file, context) pairs.
references: dict[str, list[tuple[str, str]]] = defaultdict(list)
for folder in args.wildcards:
if not folder.is_dir():
print(f"Warning: wildcard folder does not exist: {folder}", file=sys.stderr)
continue
for path in sorted(folder.rglob("*")):
if path.is_file() and path.suffix.lower() in WILDCARD_EXTENSIONS:
for name, ctx in scan_wildcard_file(path):
references[name].append((str(path), ctx))
for folder in args.enmappings:
if not folder.is_dir():
print(f"Warning: enmapping folder does not exist: {folder}", file=sys.stderr)
continue
for path in sorted(folder.rglob("*")):
if path.is_file() and path.suffix.lower() in ENMAPPING_EXTENSIONS:
for name, ctx in scan_enmapping_file(path):
references[name].append((str(path), ctx))
found_count = 0
missing_count = 0
dynamic_count = 0
for lora_name in sorted(references):
if _is_dynamic(lora_name):
dynamic_count += 1
print(f"SKIPPED (dynamic): {lora_name!r}")
for source_file, location in references[lora_name]:
print(f" {source_file} [{location}]")
continue
lookup = lora_name if args.case_sensitive else lora_name.lower()
if lookup in lora_index:
found_count += 1
if args.verbose:
print(f"OK: {lora_name}")
else:
missing_count += 1
print(f"MISSING: {lora_name}")
for source_file, location in references[lora_name]:
print(f" {source_file} [{location}]")
total = found_count + missing_count
parts = [f"{total} LoRA reference(s) checked", f"{missing_count} missing"]
if dynamic_count:
parts.append(f"{dynamic_count} dynamic (skipped)")
print(f"\n{', '.join(parts)}.")
sys.exit(1 if missing_count > 0 else 0)
if __name__ == "__main__":
main()
+15
View File
@@ -0,0 +1,15 @@
# ACB PPP Run Mode Options node
Provides run mode options to the main PPP node.
## Inputs
* **results_limit**: Limit for the number of generated results (except in `single` mode). Important for combinatorial mode.
* **results_shuffle**: It shuffles the results.
* **comb_random_fixed**: If True all specified random samplers will have a fixed value across the combinations.
* **default_sampler**: The default choice sampler when not specified (in non combinatorial mode). Also applies to extranetwork mapping selection.
* **next_seed**: Choose what to do with the seed in the following prompts in `multiple` or `combinatorial` mode. Value can be: `randomize`, `input`, `increment`, `decrement`.
## Outputs
* **options**: The options to send to the PPP node.
+3 -4
View File
@@ -15,14 +15,13 @@ Main PPP node that processes prompts.
* **process_wildcards**: Activates the wildcard processing. * **process_wildcards**: Activates the wildcard processing.
* **do_cleanup**: Activates the cleanup processing. * **do_cleanup**: Activates the cleanup processing.
* **cleanup_variables**: Do a cleanup of the output variables (depends on do_cleanup). * **cleanup_variables**: Do a cleanup of the output variables (depends on do_cleanup).
* **do_combinatorial**: Activates combinatorial mode, where the output are all the combinations of choices/wildcards of the prompt. * **results_file**: Filename to save processing results. Supports `%datetime%`, `%date%`, `%time%`, and `%host%` tokens. The file extension determines the format: `.yaml`/`.yml`, `.jsonl`, `.csv`, or plain text for any other extension. Relative paths are resolved against the extension's `logs` folder. Leave empty to disable.
* **combinatorial_shuffle**: It shuffles the combinatorial results. * **run_mode**: Sets how the process works. `single` or `multiple` for regular one or more results, or `combinatorial` for combinatorial mode, where the output are all the combinations of choices/wildcards of the prompt.
* **combinatorial_limit**: Limit for the number of generated combinations.
* **wc_options**: Connection to a Wildcards options node. * **wc_options**: Connection to a Wildcards options node.
* **stn_options**: Connection to a Send-To-Negative options node. * **stn_options**: Connection to a Send-To-Negative options node.
* **cup_options**: Connection to a Cleanup options node. * **cup_options**: Connection to a Cleanup options node.
* **en_options**: Connection to a ExtraNetworkMapping options node. * **en_options**: Connection to a ExtraNetworkMapping options node.
* **results_file**: Filename to save processing results. Supports `%datetime%`, `%date%`, `%time%`, and `%host%` tokens. The file extension determines the format: `.yaml`/`.yml`, `.jsonl`, `.csv`, or plain text for any other extension. Relative paths are resolved against the extension's `logs` folder. Leave empty to disable. * **rm_options**: Connection to a Run Mode options node.
The options nodes are optional. If you don't need to change any of the default values then you don't need to use them. The options nodes are optional. If you don't need to change any of the default values then you don't need to use them.